From e05ad27590a4ee755a8d16ba89dbd68d1320228c Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 22 Jul 2026 15:49:26 +0800 Subject: [PATCH 01/72] feat: imbalance statistics --- lightllm/common/basemodel/basemodel.py | 13 +- .../meta_weights/fused_moe/ep_balance.py | 14 + .../fused_moe/impl/deepgemm_impl.py | 24 +- .../fused_moe/grouped_fused_moe_ep.py | 10 + lightllm/distributed/communication_op.py | 9 + lightllm/server/api_cli.py | 5 + lightllm/server/core/objs/start_args_type.py | 1 + lightllm/server/metrics/metrics.py | 12 + .../mode_backend/ep_balance_monitor.py | 346 +++++++++++ .../server/router/model_infer/model_rpc.py | 9 + .../model_infer/test_ep_balance_monitor.py | 546 ++++++++++++++++++ 11 files changed, 984 insertions(+), 5 deletions(-) create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py create mode 100644 lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py create mode 100644 unit_tests/server/router/model_infer/test_ep_balance_monitor.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index f59535cb13..b5d364e61a 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -70,6 +70,7 @@ class TpPartBaseModel: def __init__(self, kvargs): self.args = get_env_start_args() + self.ep_balance_monitor = None self.run_mode = kvargs["run_mode"] self.weight_dir_ = kvargs["weight_dir"] self.max_total_token_num = kvargs["max_total_token_num"] @@ -316,9 +317,14 @@ def forward(self, model_input: ModelInput): assert model_input.mem_indexes.is_cuda if model_input.is_prefill: - return self._prefill(model_input=model_input) - else: - return self._decode(model_input) + model_output = self._prefill(model_input=model_input) + self._record_prefill_ep_balance() + return model_output + return self._decode(model_input) + + def _record_prefill_ep_balance(self): + if self.ep_balance_monitor is not None: + self.ep_balance_monitor.record_prefill_round() def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() @@ -814,6 +820,7 @@ def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input dist_group_manager.clear_deepep_buffer() model_output0.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event model_output1.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event + self._record_prefill_ep_balance() return model_output0, model_output1 @torch.no_grad() diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py new file mode 100644 index 0000000000..70435146c5 --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py @@ -0,0 +1,14 @@ +from dataclasses import dataclass + + +@dataclass(slots=True) +class PrefillEPBalanceCounters: + """Cumulative CPU loads for one EP MoE layer's completed prefill dispatches.""" + + route_load: int = 0 + compute_load: int = 0 + + def accumulate(self, route_load: int, compute_load: int): + """Accumulate exact route and alignment-expanded compute loads for one prefill dispatch.""" + self.route_load += route_load + self.compute_load += compute_load 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 024be9f55c..ec2c1fd39c 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 @@ -20,6 +20,10 @@ class FuseMoeDeepGEMM(FuseMoeTriton): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.ep_balance_counters = None + def _select_experts( self, input_tensor: torch.Tensor, @@ -87,6 +91,7 @@ def _fused_experts( quant_method=self.quant_method, is_prefill=is_prefill, previous_event=None, # for overlap + ep_balance_counters=self.ep_balance_counters, ) return output @@ -181,8 +186,23 @@ def dispatch( use_tma_aligned_col_major_sf=True, ) - def hook(): - event.current_stream_wait() + counters = self.ep_balance_counters + if counters is None: + + def hook(): + event.current_stream_wait() + + else: + # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. + route_load = topk_idx.numel() + compute_load = recv_x[0].shape[0] + + def hook(): + event.current_stream_wait() + counters.accumulate( + route_load=route_load, + compute_load=compute_load, + ) return recv_x, recv_topk_idx, recv_topk_weights, handle.num_recv_tokens_per_expert_list, handle, hook 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 ca39376bab..ff430e5e51 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 @@ -19,6 +19,7 @@ ep_gather_chunk, ep_zero_padding, ) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters from lightllm.utils.envs_utils import ( get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, @@ -201,6 +202,7 @@ def fused_experts( quant_method: Any, is_prefill: Optional[bool], previous_event: Optional[Any] = None, + ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): check_ep_expert_dtype(quant_method) if use_sm100_mega_moe(quant_method): @@ -222,6 +224,7 @@ def fused_experts( w1_scale=w13.weight_scale, w2_scale=w2.weight_scale, previous_event=previous_event, + ep_balance_counters=ep_balance_counters, ) @@ -240,6 +243,7 @@ def fused_experts_impl( w1_scale: Optional[torch.Tensor] = None, w2_scale: Optional[torch.Tensor] = None, previous_event: Optional[Any] = None, + ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): # Check constraints. assert hidden_states.shape[1] == w1.shape[2], "Hidden size mismatch" @@ -290,6 +294,12 @@ def fused_experts_impl( do_expand=True, use_tma_aligned_col_major_sf=True, ) + if ep_balance_counters is not None: + # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. + ep_balance_counters.accumulate( + route_load=topk_idx.numel(), + compute_load=recv_x[0].shape[0], + ) # 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 diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 93c603212d..76b202f4d5 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -108,6 +108,7 @@ def all_gather_into_tensor(self, output_: torch.Tensor, input_: torch.Tensor, as class DistributeGroupManager: def __init__(self): self.groups = [] + self.ep_balance_monitor_group = None self.ep_buffer = None self.ep_low_latency_buffer = None self.ep_mega_moe_buffer = None @@ -125,6 +126,14 @@ def create_groups(self, group_size: int): if not args.disable_flashinfer_allreduce: group.init_flashinfer_reduce() self.groups.append(group) + if ( + getattr(args, "enable_ep_moe", False) + and not getattr(args, "disable_ep_balance_monitor", False) + and getattr(args, "run_mode", "normal") != "decode" + and not getattr(args, "enable_prefill_cudagraph", False) + and not is_sm100_gpu() + ): + self.ep_balance_monitor_group = dist.new_group(ranks=list(range(get_global_world_size())), backend="gloo") return def get_default_group(self) -> CustomProcessGroup: diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index e01789d946..93656f600b 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -766,6 +766,11 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="""Whether to enable ep moe for deepseekv3 model.""", ) + parser.add_argument( + "--disable_ep_balance_monitor", + action="store_true", + help="""Disable the prefill expert balance monitor enabled by default for EP-MoE.""", + ) parser.add_argument( "--ep_redundancy_expert_config_path", type=str, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 254099d5bc..16ee99b3c7 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -186,6 +186,7 @@ class StartArgs: default="gpu_counter", metadata={"choices": ["cpu_counter", "pin_mem_counter", "gpu_counter"]} ) enable_ep_moe: bool = field(default=False) + disable_ep_balance_monitor: bool = field(default=False) ep_redundancy_expert_config_path: Optional[str] = field(default=None) auto_update_redundancy_expert: bool = field(default=False) enable_fused_shared_experts: bool = field(default=False) diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index 0d42462c3f..c19d756c17 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -32,6 +32,15 @@ "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", "lightllm_gen_throughput": "Generation throughput of latest completed request (tokens/s)", "lightllm_num_running_reqs": "Number of running requests", + "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token": ( + "Estimated critical-path excess compute per logical routed source token, GFLOPs/token" + ), + "lightllm_prefill_ep_compute_critical_overhead_ratio": ( + "Estimated excess critical compute divided by balanced compute; 0.3 means +30%" + ), + "lightllm_prefill_ep_placement_pressure_drift": ( + "Normalized temporal drift of overloaded-rank pressure from the latest complete prefill report" + ), } @@ -111,6 +120,9 @@ def init_metrics(self, args): self.create_gauge("lightllm_cache_hit_rate") self.create_gauge("lightllm_gen_throughput") self.create_gauge("lightllm_num_running_reqs") + self.create_gauge("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token") + self.create_gauge("lightllm_prefill_ep_compute_critical_overhead_ratio") + self.create_gauge("lightllm_prefill_ep_placement_pressure_drift") def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) diff --git a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py new file mode 100644 index 0000000000..107f4ba504 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py @@ -0,0 +1,346 @@ +import threading +from array import array +from typing import Optional, Tuple + +import torch +import torch.distributed as dist + +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters +from lightllm.distributed.communication_op import dist_group_manager +from lightllm.server.metrics.manager import MetricClient +from lightllm.utils.device_utils import is_sm100_gpu +from lightllm.utils.dist_utils import ( + get_global_rank, + get_global_world_size, +) +from lightllm.utils.log_utils import init_logger +from lightllm.utils.shm_port_args import get_shm_port_args + + +logger = init_logger(__name__) + +EP_BALANCE_PREFILL_ROUNDS_PER_REPORT = 100 +EP_BALANCE_ROUND_BUFFER_CAPACITY = 4096 +EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS = 20 +EP_BALANCE_PRESSURE_DRIFT_STABLE_THRESHOLD = 0.10 +ROUTE_LOAD = 0 +COMPUTE_LOAD = 1 +GFLOP = 1_000_000_000 + + +def should_enable_ep_balance_monitor(args) -> bool: + if args.enable_prefill_cudagraph or is_sm100_gpu(): + return False + return args.enable_ep_moe and not args.disable_ep_balance_monitor and args.run_mode != "decode" + + +def calculate_prefill_balance_stats( + round_stats: torch.Tensor, # [num_rounds, num_layers, world_size, 2] (route/compute) + layer_routed_experts: torch.Tensor, # [num_layers] + layer_flops_per_expert_token: torch.Tensor, # [num_layers] + layer_topks: torch.Tensor, # [num_layers] + source_token_replication: int, + report_min_route_samples_per_expert: int = 100, +) -> Optional[dict]: + """Summarize complete-prefill samples from [round, layer, rank, route/compute]. + + MoE layers execute sequentially, and every layer waits for its slowest EP + rank. Preserve the layer dimension until after taking the cross-rank max so + that different slow ranks in different layers cannot cancel each other. + """ + assert source_token_replication > 0 + + layer_route_load = round_stats[:, :, :, ROUTE_LOAD].sum(dim=(0, 2)) + minimum_route_samples = layer_routed_experts * report_min_route_samples_per_expert + if torch.any(layer_route_load < minimum_route_samples): + return None + total_route_load = layer_route_load.sum() + + compute_rank_load = round_stats[:, :, :, COMPUTE_LOAD] + compute_rank_load_float = compute_rank_load.to(torch.float64) + excess_compute_load = compute_rank_load_float.max(dim=2).values - compute_rank_load_float.mean(dim=2) + + # Every MoE expert token executes the two projections packed in w13 plus + # the w2 projection. Weight each padded compute token by the layer's + # actual matrix sizes so the metric remains comparable across models. + excess_compute_flops = (excess_compute_load * layer_flops_per_expert_token.to(torch.float64)).sum() + balanced_compute_flops = ( + compute_rank_load_float.mean(dim=2) * layer_flops_per_expert_token.to(torch.float64) + ).sum() + if balanced_compute_flops == 0: + return None + + # Non-TPSP prefill gathers one route-load copy per TP rank. Divide the + # replica count out so GFLOP/token uses logical source tokens. + source_tokens = total_route_load.to(torch.float64) / ( + layer_topks.to(torch.float64).sum() * source_token_replication + ) + if source_tokens == 0: + return None + + return { + "prefill_rounds": int(compute_rank_load.shape[0]), + "critical_overhead_gflops_per_routed_token": float((excess_compute_flops / source_tokens / GFLOP).item()), + "prefill_ep_compute_critical_overhead_ratio": float((excess_compute_flops / balanced_compute_flops).item()), + } + + +def calculate_prefill_placement_pressure_drift( + round_stats: torch.Tensor, + previous_pressure_signature: Optional[torch.Tensor] = None, + bucket_rounds: int = EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS, +) -> Tuple[float, torch.Tensor]: + """Measure how prefill rank-pressure placement changes across time buckets. + + This is rank-0 CPU-only report analysis. The returned final bucket is a + compact signature that allows the next report to include the boundary pair. + """ + if bucket_rounds <= 0: + raise ValueError(f"bucket_rounds must be positive, got {bucket_rounds}") + if round_stats.ndim != 4 or round_stats.shape[-1] != 2: + raise ValueError( + "round_stats must have shape [num_rounds, num_layers, world_size, 2], " f"got {tuple(round_stats.shape)}" + ) + num_rounds, num_layers, world_size, _ = round_stats.shape + if num_rounds <= 0: + raise ValueError("round_stats must contain at least one round") + if num_layers <= 0 or world_size <= 0: + raise ValueError("round_stats must contain at least one layer and rank") + if num_rounds % bucket_rounds != 0: + raise ValueError(f"num_rounds ({num_rounds}) must be divisible by bucket_rounds ({bucket_rounds})") + + num_buckets = num_rounds // bucket_rounds + bucket_rank_load = ( + round_stats[:, :, :, COMPUTE_LOAD] + .to(torch.float64) + .reshape(num_buckets, bucket_rounds, num_layers, world_size) + .sum(dim=1) + ) + mean_rank_load = bucket_rank_load.mean(dim=2, keepdim=True).clamp_min(1) + pressure = torch.relu(bucket_rank_load / mean_rank_load - 1) + + if previous_pressure_signature is not None: + expected_shape = (num_layers, world_size) + if tuple(previous_pressure_signature.shape) != expected_shape: + raise ValueError( + "previous_pressure_signature must have shape " + f"{expected_shape}, got {tuple(previous_pressure_signature.shape)}" + ) + left = torch.cat((previous_pressure_signature.to(torch.float64).unsqueeze(0), pressure[:-1]), dim=0) + right = pressure + else: + left = pressure[:-1] + right = pressure[1:] + + total_pressure = (left + right).sum() + if total_pressure == 0: + drift = 0.0 + else: + drift = float((left - right).abs().sum().div(total_pressure).item()) + return drift, pressure[-1].clone() + + +def classify_prefill_placement_pressure_drift(drift: float) -> str: + if drift < EP_BALANCE_PRESSURE_DRIFT_STABLE_THRESHOLD: + return "stable" + return "dynamic" + + +def _find_fused_moe_weights(model): + weights_by_id = {} + for layer in model.trans_layers_weight: + for value in getattr(layer, "__dict__", {}).values(): + if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: + weights_by_id[id(value)] = value + return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) + + +class EPBalanceMonitor: + """Report cross-rank imbalance for non-overlapping blocks of complete prefill rounds.""" + + def __init__(self, model: TpPartBaseModel): + self.global_rank = get_global_rank() + self.world_size = get_global_world_size() + self.weights = _find_fused_moe_weights(model) + self.enabled = bool(self.weights) + + if not self.enabled: + return + + self.source_token_replication = 1 if model.args.enable_tpsp_mix_mode else model.tp_world_size_ + self.counters: list[PrefillEPBalanceCounters] = [PrefillEPBalanceCounters() for _ in self.weights] + for weight, counter in zip(self.weights, self.counters): + weight.fuse_moe_impl.ep_balance_counters = counter + self.layer_routed_experts = torch.tensor( + [weight.n_routed_experts for weight in self.weights], dtype=torch.int64 + ) + self.layer_flops_per_expert_token = torch.tensor( + [ + # Each expert-token performs gate, up, and down projections; each MAC counts as 2 FLOPs. + 2 * 3 * weight.hidden_size * weight.moe_intermediate_size + for weight in self.weights + ], + dtype=torch.float64, + ) + self.layer_topks = torch.tensor( + [weight.num_experts_per_tok for weight in self.weights], + dtype=torch.float64, + ) + self._round_buffer_storage = array("q", [0]) * (EP_BALANCE_ROUND_BUFFER_CAPACITY * len(self.weights) * 2) + self._round_buffer = torch.frombuffer(self._round_buffer_storage, dtype=torch.int64).view( + EP_BALANCE_ROUND_BUFFER_CAPACITY, len(self.weights), 2 + ) + self._round_ready = threading.Event() + self._written_round_count = 0 # Prefill rounds fully written to the ring buffer. + self._processed_round_count = 0 # Prefill rounds consumed by the monitor thread. + self._overflowed = False + self._previous_pressure_signature: Optional[torch.Tensor] = None + self._common_round_end = torch.zeros((), dtype=torch.int64) + + self.gloo_group = dist_group_manager.ep_balance_monitor_group + if self.gloo_group is None: + raise RuntimeError("EP balance monitor requires a pre-created dedicated Gloo process group") + self.metric_client = MetricClient(get_shm_port_args().metric_port) if self.global_rank == 0 else None + threading.Thread(target=self._monitor_loop, daemon=True, name="ep-balance-monitor").start() + + def record_prefill_round(self): + """Publish one complete all-layer prefill sample to the SPSC ring.""" + if not self.enabled: + return + + written_round_count = self._written_round_count + if written_round_count - self._processed_round_count >= EP_BALANCE_ROUND_BUFFER_CAPACITY: + if not self._overflowed: + self._overflowed = True + self._round_ready.set() + return + + storage_index = (written_round_count % EP_BALANCE_ROUND_BUFFER_CAPACITY) * len(self.counters) * 2 + for counter in self.counters: + self._round_buffer_storage[storage_index] = counter.route_load + self._round_buffer_storage[storage_index + 1] = counter.compute_load + counter.route_load = 0 + counter.compute_load = 0 + storage_index += 2 + + # Publish only after the entire slot is written. The SPSC producer and + # monitor thread run under the CPython GIL, so this count is the release + # point for the corresponding ring slot. + self._written_round_count = written_round_count + 1 + if self._written_round_count - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: + self._round_ready.set() + + def _get_common_round_end(self) -> int: + """Return the exclusive round boundary completed by every rank.""" + self._common_round_end.fill_(self._written_round_count) + dist.all_reduce(self._common_round_end, op=dist.ReduceOp.MIN, group=self.gloo_group) + return int(self._common_round_end.item()) + + def _raise_buffer_overflow(self, phase: str, common_round_end: Optional[int] = None): + message = ( + "EP balance prefill-round buffer overflowed " + f"phase={phase} written={self._written_round_count} " + f"processed={self._processed_round_count} capacity={EP_BALANCE_ROUND_BUFFER_CAPACITY}" + ) + if common_round_end is not None: + message += f" common_round_end={common_round_end}" + raise RuntimeError(message) + + def _copy_local_rounds(self, start: int, end: int) -> torch.Tensor: + """Copy local prefill-round loads in the half-open range [start, end).""" + num_rounds = end - start + if num_rounds > EP_BALANCE_ROUND_BUFFER_CAPACITY: + raise ValueError("requested EP balance round range exceeds ring capacity") + start_index = start % EP_BALANCE_ROUND_BUFFER_CAPACITY + if start_index + num_rounds <= EP_BALANCE_ROUND_BUFFER_CAPACITY: + return self._round_buffer[start_index : start_index + num_rounds].clone() + end_index = (start_index + num_rounds) % EP_BALANCE_ROUND_BUFFER_CAPACITY + return torch.cat((self._round_buffer[start_index:], self._round_buffer[:end_index]), dim=0) + + def _gather_round_stats(self, local_round_stats: torch.Tensor) -> Optional[torch.Tensor]: + """Gather rank-local stats as [round, layer, rank, route/compute].""" + gathered = ( + [torch.empty_like(local_round_stats) for _ in range(self.world_size)] if self.global_rank == 0 else None + ) + dist.gather(local_round_stats, gather_list=gathered, dst=0, group=self.gloo_group) + if self.global_rank != 0: + return None + # [rank, round, layer, route/compute] + # -> [round, layer, rank, route/compute] + return torch.stack(gathered).permute(1, 2, 0, 3) + + def _log_stats(self, round_stats: torch.Tensor): + """Compute and log balance statistics for one complete global window.""" + compute = calculate_prefill_balance_stats( + round_stats, + self.layer_routed_experts, + self.layer_flops_per_expert_token, + self.layer_topks, + self.source_token_replication, + ) + if compute is None: + return + + drift, self._previous_pressure_signature = calculate_prefill_placement_pressure_drift( + round_stats, + previous_pressure_signature=self._previous_pressure_signature, + ) + drift_state = classify_prefill_placement_pressure_drift(drift) + + logger.info( + "ep_balance " + f"phase=prefill prefill_rounds={compute['prefill_rounds']} " + "prefill_ep_critical_overhead_gflops_per_routed_token=" + f"{compute['critical_overhead_gflops_per_routed_token']:.4f} " + "prefill_ep_compute_critical_overhead_ratio=" + f"{compute['prefill_ep_compute_critical_overhead_ratio']:.4f} " + f"prefill_ep_placement_pressure_drift={drift:.4f} " + f"prefill_ep_placement_pressure_state={drift_state}" + ) + self.metric_client.gauge_set( + "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", + compute["critical_overhead_gflops_per_routed_token"], + ) + self.metric_client.gauge_set( + "lightllm_prefill_ep_compute_critical_overhead_ratio", + compute["prefill_ep_compute_critical_overhead_ratio"], + ) + self.metric_client.gauge_set("lightllm_prefill_ep_placement_pressure_drift", drift) + + def _monitor_loop(self): + """Consume commonly completed rounds in background report-sized windows.""" + try: + while True: + self._round_ready.wait() + self._round_ready.clear() + if self._overflowed: + self._raise_buffer_overflow("before_sync") + common_round_end = self._get_common_round_end() + if common_round_end - self._processed_round_count > EP_BALANCE_ROUND_BUFFER_CAPACITY: + self._raise_buffer_overflow("common_round_lag", common_round_end=common_round_end) + + while common_round_end - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: + round_start = self._processed_round_count + round_end = round_start + EP_BALANCE_PREFILL_ROUNDS_PER_REPORT + local_round_stats = self._copy_local_rounds(round_start, round_end) + if self._overflowed: + self._raise_buffer_overflow("after_copy") + round_stats = self._gather_round_stats(local_round_stats) + self._processed_round_count = round_end + if self.global_rank == 0: + self._log_stats(round_stats) + + if self._written_round_count - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: + self._round_ready.set() + except Exception as exc: + logger.exception(f"EP balance monitor stopped unexpectedly: {exc}") + self._disable() + return + + def _disable(self): + """Detach counters from MoE weights and disable monitoring.""" + for weight in self.weights: + weight.fuse_moe_impl.ep_balance_counters = None + self.enabled = False diff --git a/lightllm/server/router/model_infer/model_rpc.py b/lightllm/server/router/model_infer/model_rpc.py index 17dd96be85..3ae4d4cbc2 100644 --- a/lightllm/server/router/model_infer/model_rpc.py +++ b/lightllm/server/router/model_infer/model_rpc.py @@ -27,6 +27,10 @@ ) from lightllm.server.router.model_infer.mode_backend.redundancy_expert_manager import RedundancyExpertManager from lightllm.server.router.model_infer.mode_backend.rl_backend_ops import RlBackendOps +from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( + EPBalanceMonitor, + should_enable_ep_balance_monitor, +) from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.utils.log_utils import init_logger from lightllm.utils.graceful_utils import graceful_registry @@ -104,6 +108,11 @@ def exposed_init_model(self, kvargs): logger.info("init redundancy_expert_manager") else: self.redundancy_expert_manager = None + + if should_enable_ep_balance_monitor(self.args): + monitor = EPBalanceMonitor(self.backend.model) + if monitor.enabled: + self.backend.model.ep_balance_monitor = monitor return def exposed_get_max_total_token_num(self): diff --git a/unit_tests/server/router/model_infer/test_ep_balance_monitor.py b/unit_tests/server/router/model_infer/test_ep_balance_monitor.py new file mode 100644 index 0000000000..2ddcaafddf --- /dev/null +++ b/unit_tests/server/router/model_infer/test_ep_balance_monitor.py @@ -0,0 +1,546 @@ +import threading +from array import array +from types import SimpleNamespace + +import pytest +import torch +from prometheus_client import generate_latest + +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters +from lightllm.distributed import communication_op as communication_op_module +from lightllm.server.metrics.metrics import Monitor +from lightllm.server.router.model_infer.mode_backend import ep_balance_monitor as monitor_module +from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( + EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS, + calculate_prefill_placement_pressure_drift, + calculate_prefill_balance_stats, + classify_prefill_placement_pressure_drift, + should_enable_ep_balance_monitor, +) + + +@pytest.fixture(autouse=True) +def _mock_non_sm100(monkeypatch): + monkeypatch.setattr(monitor_module, "is_sm100_gpu", lambda: False) + + +def _stats(source_token_replication: int): + return calculate_prefill_balance_stats( + torch.tensor([[[[800, 40], [800, 20]], [[800, 20], [800, 40]]]], dtype=torch.int64), + layer_routed_experts=torch.tensor([1, 1]), + layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), + layer_topks=torch.tensor([2.0, 2.0]), + source_token_replication=source_token_replication, + ) + + +def _monitor_args(**overrides): + args = { + "enable_ep_moe": True, + "disable_ep_balance_monitor": False, + "run_mode": "normal", + "enable_prefill_cudagraph": False, + } + args.update(overrides) + return SimpleNamespace(**args) + + +def _pressure_round_stats(bucket_rank_loads): + """Build [round, layer=1, rank, route/compute] CPU samples for drift tests.""" + return torch.tensor([[[[0, load] for load in rank_loads]] for rank_loads in bucket_rank_loads], dtype=torch.int64) + + +def test_pressure_drift_is_zero_for_identical_pressure(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [2, 1]]), bucket_rounds=1) + assert drift == 0.0 + + +def test_pressure_drift_is_one_for_complete_hot_rank_migration(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 0], [0, 2]]), bucket_rounds=1) + assert drift == 1.0 + + +def test_pressure_drift_tracks_same_rank_magnitude_change(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [3, 1]]), bucket_rounds=1) + assert drift == pytest.approx(0.2) + + +def test_pressure_drift_is_invariant_to_uniform_load_scale(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [4, 2]]), bucket_rounds=1) + assert drift == 0.0 + + +def test_pressure_drift_is_zero_for_balanced_inputs(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[8, 8], [16, 16]]), bucket_rounds=1) + assert drift == 0.0 + + +def test_pressure_drift_previous_signature_bridges_report_boundary(): + _, signature = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 0]]), bucket_rounds=1) + drift, next_signature = calculate_prefill_placement_pressure_drift( + _pressure_round_stats([[0, 2]]), previous_pressure_signature=signature, bucket_rounds=1 + ) + assert drift == 1.0 + assert torch.equal(next_signature, torch.tensor([[0.0, 1.0]], dtype=torch.float64)) + + +def test_pressure_drift_default_bucket_handles_normal_report_window(): + round_stats = _pressure_round_stats([[4, 2]] * 100) + drift, signature = calculate_prefill_placement_pressure_drift(round_stats) + assert EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS == 20 + assert drift == 0.0 + assert signature.shape == (1, 2) + + +@pytest.mark.parametrize( + ("drift", "expected"), + [ + (0.0, "stable"), + (0.0999, "stable"), + (0.10, "dynamic"), + (0.2999, "dynamic"), + (0.30, "dynamic"), + (0.75, "dynamic"), + (1.0, "dynamic"), + ], +) +def test_pressure_drift_classification_boundaries(drift, expected): + assert classify_prefill_placement_pressure_drift(drift) == expected + + +def test_monitor_log_stats_reports_pressure_drift_and_bridges_reports(monkeypatch): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.layer_routed_experts = torch.tensor([1], dtype=torch.int64) + monitor.layer_flops_per_expert_token = torch.tensor([2.0], dtype=torch.float64) + monitor.layer_topks = torch.tensor([2.0], dtype=torch.float64) + monitor.source_token_replication = 1 + monitor._previous_pressure_signature = None + metric_calls = [] + monitor.metric_client = SimpleNamespace(gauge_set=lambda name, value: metric_calls.append((name, value))) + logs = [] + monkeypatch.setattr(monitor_module.logger, "info", logs.append) + + def report_for_hot_rank(hot_rank): + round_stats = torch.zeros((100, 1, 2, 2), dtype=torch.int64) + round_stats[:, :, :, monitor_module.ROUTE_LOAD] = 100 + round_stats[:, :, hot_rank, monitor_module.COMPUTE_LOAD] = 2 + monitor._log_stats(round_stats) + + report_for_hot_rank(0) + first_signature = monitor._previous_pressure_signature.clone() + report_for_hot_rank(1) + + assert "prefill_ep_placement_pressure_drift=0.0000" in logs[0] + assert "prefill_ep_placement_pressure_state=stable" in logs[0] + assert "prefill_ep_placement_pressure_drift=0.2000" in logs[1] + assert "prefill_ep_placement_pressure_state=dynamic" in logs[1] + assert torch.equal(first_signature, torch.tensor([[1.0, 0.0]], dtype=torch.float64)) + assert torch.equal(monitor._previous_pressure_signature, torch.tensor([[0.0, 1.0]], dtype=torch.float64)) + assert metric_calls == [ + ("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", pytest.approx(2e-11)), + ("lightllm_prefill_ep_compute_critical_overhead_ratio", pytest.approx(1.0)), + ("lightllm_prefill_ep_placement_pressure_drift", 0.0), + ("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", pytest.approx(2e-11)), + ("lightllm_prefill_ep_compute_critical_overhead_ratio", pytest.approx(1.0)), + ("lightllm_prefill_ep_placement_pressure_drift", pytest.approx(0.2)), + ] + + +def test_critical_overhead_preserves_per_layer_slowest_rank(): + stats = _stats(source_token_replication=1) + assert stats is not None + assert stats["critical_overhead_gflops_per_routed_token"] == pytest.approx(7.5e-11) + assert stats["prefill_ep_compute_critical_overhead_ratio"] == pytest.approx(1 / 3) + + +def test_non_tpsp_tp_replication_scales_gflops_per_token_but_not_ratio(): + tpsp_stats = _stats(source_token_replication=1) + non_tpsp_tp8_stats = _stats(source_token_replication=8) + assert tpsp_stats is not None and non_tpsp_tp8_stats is not None + assert non_tpsp_tp8_stats["critical_overhead_gflops_per_routed_token"] == pytest.approx( + tpsp_stats["critical_overhead_gflops_per_routed_token"] * 8 + ) + assert non_tpsp_tp8_stats["prefill_ep_compute_critical_overhead_ratio"] == pytest.approx( + tpsp_stats["prefill_ep_compute_critical_overhead_ratio"] + ) + + +def test_critical_overhead_is_zero_when_ranks_are_balanced(): + stats = calculate_prefill_balance_stats( + torch.tensor([[[[100, 32], [100, 32]], [[100, 64], [100, 64]]]], dtype=torch.int64), + layer_routed_experts=torch.tensor([1, 1]), + layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), + layer_topks=torch.tensor([2.0, 2.0]), + source_token_replication=1, + ) + assert stats is not None + assert stats["critical_overhead_gflops_per_routed_token"] == 0.0 + assert stats["prefill_ep_compute_critical_overhead_ratio"] == 0.0 + + +@pytest.mark.parametrize( + "round_stats", + [ + torch.tensor([[[[1, 1], [1, 1]]]], dtype=torch.int64), + torch.tensor([[[[100, 0], [100, 0]]]], dtype=torch.int64), + ], +) +def test_critical_overhead_rejects_insufficient_or_zero_compute_samples(round_stats): + assert ( + calculate_prefill_balance_stats( + round_stats, + layer_routed_experts=torch.tensor([1]), + layer_flops_per_expert_token=torch.tensor([2.0]), + layer_topks=torch.tensor([2.0]), + source_token_replication=1, + ) + is None + ) + + +def test_cpu_counter_accumulates_multiple_prefill_dispatches(): + counters = PrefillEPBalanceCounters() + counters.accumulate(route_load=3, compute_load=128) + counters.accumulate(route_load=4, compute_load=256) + assert (counters.route_load, counters.compute_load) == (7, 384) + + +def test_monitor_reuses_manager_precreated_dedicated_gloo_group(monkeypatch): + sentinel_group = object() + impl = SimpleNamespace(ep_balance_counters=None) + weight = SimpleNamespace( + fuse_moe_impl=impl, + n_routed_experts=8, + hidden_size=16, + moe_intermediate_size=32, + num_experts_per_tok=2, + ) + model = SimpleNamespace(args=SimpleNamespace(enable_tpsp_mix_mode=True), tp_world_size_=1) + + monkeypatch.setattr(monitor_module, "_find_fused_moe_weights", lambda _: [weight]) + monkeypatch.setattr(monitor_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(monitor_module, "get_global_world_size", lambda: 2) + metric_ports = [] + monkeypatch.setattr(monitor_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=4321)) + monkeypatch.setattr(monitor_module, "MetricClient", lambda port: metric_ports.append(port) or object()) + monkeypatch.setattr(monitor_module, "dist_group_manager", SimpleNamespace(ep_balance_monitor_group=sentinel_group)) + monkeypatch.setattr(monitor_module.dist, "new_group", lambda *args, **kwargs: pytest.fail("unexpected new_group")) + monkeypatch.setattr( + monitor_module.threading, + "Thread", + lambda *args, **kwargs: SimpleNamespace(start=lambda: None), + ) + + monitor = monitor_module.EPBalanceMonitor(model) + + assert monitor.gloo_group is sentinel_group + assert impl.ep_balance_counters is monitor.counters[0] + assert metric_ports == [4321] + + +def test_nonzero_rank_monitor_does_not_create_metric_client(monkeypatch): + sentinel_group = object() + impl = SimpleNamespace(ep_balance_counters=None) + weight = SimpleNamespace( + fuse_moe_impl=impl, + n_routed_experts=8, + hidden_size=16, + moe_intermediate_size=32, + num_experts_per_tok=2, + ) + model = SimpleNamespace(args=SimpleNamespace(enable_tpsp_mix_mode=True), tp_world_size_=1) + metric_client_calls = [] + monkeypatch.setattr(monitor_module, "_find_fused_moe_weights", lambda _: [weight]) + monkeypatch.setattr(monitor_module, "get_global_rank", lambda: 1) + monkeypatch.setattr(monitor_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(monitor_module, "dist_group_manager", SimpleNamespace(ep_balance_monitor_group=sentinel_group)) + monkeypatch.setattr(monitor_module.dist, "new_group", lambda *args, **kwargs: pytest.fail("unexpected new_group")) + monkeypatch.setattr(monitor_module, "get_shm_port_args", lambda: pytest.fail("unexpected port lookup")) + monkeypatch.setattr( + monitor_module, + "MetricClient", + lambda port: metric_client_calls.append(port) or pytest.fail("unexpected metric client"), + ) + monkeypatch.setattr( + monitor_module.threading, + "Thread", + lambda *args, **kwargs: SimpleNamespace(start=lambda: None), + ) + + monitor = monitor_module.EPBalanceMonitor(model) + + assert monitor.metric_client is None + assert metric_client_calls == [] + + +@pytest.mark.parametrize("disable_monitor", [False, True]) +def test_group_manager_creates_monitor_gloo_group_only_when_enabled(monkeypatch, disable_monitor): + monitor_group = object() + custom_groups = [] + + class FakeCustomProcessGroup: + def init_symm_mem_reduce(self): + pass + + def init_flashinfer_reduce(self): + pass + + args = SimpleNamespace( + enable_ep_moe=True, + disable_ep_balance_monitor=disable_monitor, + run_mode="normal", + enable_prefill_cudagraph=False, + disable_symm_mem_allreduce=True, + disable_flashinfer_allreduce=True, + ) + monkeypatch.setattr(communication_op_module, "get_env_start_args", lambda: args) + monkeypatch.setattr( + communication_op_module, + "CustomProcessGroup", + lambda: custom_groups.append(FakeCustomProcessGroup()) or custom_groups[-1], + ) + monkeypatch.setattr(communication_op_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(communication_op_module, "is_sm100_gpu", lambda: False) + calls = [] + monkeypatch.setattr( + communication_op_module.dist, + "new_group", + lambda *args, **kwargs: calls.append((args, kwargs)) or monitor_group, + ) + + manager = communication_op_module.DistributeGroupManager() + manager.create_groups(group_size=2) + + assert len(manager.groups) == 2 + if disable_monitor: + assert calls == [] + assert manager.ep_balance_monitor_group is None + else: + assert calls == [((), {"ranks": [0, 1], "backend": "gloo"})] + assert manager.ep_balance_monitor_group is monitor_group + + +def test_monitor_registers_prefill_ep_gauges_with_model_label(): + monitor = Monitor( + SimpleNamespace( + metric_gateway=None, + job_name="test", + grouping_key=[], + enable_monitor_auth=False, + model_name="monitor-test-model", + max_req_total_len=128, + mtp_step=0, + ) + ) + values = { + "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token": 1.25, + "lightllm_prefill_ep_compute_critical_overhead_ratio": 0.3, + "lightllm_prefill_ep_placement_pressure_drift": 0.125, + } + assert set(values).issubset(monitor.monitor_registry) + for name, value in values.items(): + monitor.gauge_set(name, value) + + exposition = generate_latest(monitor.registry).decode() + for name, value in values.items(): + assert f'{name}{{model_name="monitor-test-model"}} {value}' in exposition + + +def test_record_prefill_round_stores_cumulative_counter_deltas_in_ring_buffer(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.enabled = True + monitor.counters = [PrefillEPBalanceCounters()] + monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) + monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( + monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 + ) + monitor._round_ready = threading.Event() + monitor._written_round_count = 0 + monitor._processed_round_count = 0 + monitor._overflowed = False + + monitor.counters[0].accumulate(route_load=3, compute_load=128) + monitor.record_prefill_round() + assert (monitor.counters[0].route_load, monitor.counters[0].compute_load) == (0, 0) + monitor.counters[0].accumulate(route_load=2, compute_load=256) + monitor.record_prefill_round() + + assert torch.equal( + monitor._copy_local_rounds(0, 2), + torch.tensor([[[3, 128]], [[2, 256]]], dtype=torch.int64), + ) + assert not hasattr(monitor, "_round_lock") + + +def test_spsc_ring_copy_wraps_without_a_lock(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.enabled = True + monitor.counters = [PrefillEPBalanceCounters()] + monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) + monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( + monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 + ) + monitor._round_ready = threading.Event() + monitor._written_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2 + monitor._processed_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2 + monitor._overflowed = False + + for value in (11, 12, 13): + monitor.counters[0].accumulate(route_load=value, compute_load=value * 10) + monitor.record_prefill_round() + + assert torch.equal( + monitor._copy_local_rounds(monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2, monitor._written_round_count), + torch.tensor([[[11, 110]], [[12, 120]], [[13, 130]]], dtype=torch.int64), + ) + assert not hasattr(monitor, "_round_lock") + + +def test_spsc_ring_overflow_is_deferred_to_the_monitor_thread(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.enabled = True + monitor.counters = [PrefillEPBalanceCounters(route_load=7, compute_load=70)] + monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) + monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( + monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 + ) + monitor._round_ready = threading.Event() + monitor._written_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY + monitor._processed_round_count = 0 + monitor._overflowed = False + + monitor.record_prefill_round() + + assert monitor._overflowed + assert monitor._round_ready.is_set() + assert monitor._written_round_count == monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY + assert (monitor.counters[0].route_load, monitor.counters[0].compute_load) == (7, 70) + + +def test_raise_buffer_overflow_always_reports_phase_and_ring_counts(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor._written_round_count = 23 + monitor._processed_round_count = 7 + + with pytest.raises(RuntimeError) as exc_info: + monitor._raise_buffer_overflow("before_sync") + + assert str(exc_info.value) == ( + "EP balance prefill-round buffer overflowed " + f"phase=before_sync written=23 processed=7 capacity={monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY}" + ) + + +def test_raise_buffer_overflow_optionally_reports_common_round_end(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor._written_round_count = 23 + monitor._processed_round_count = 7 + + with pytest.raises(RuntimeError) as exc_info: + monitor._raise_buffer_overflow("common_round_lag", common_round_end=19) + + assert str(exc_info.value) == ( + "EP balance prefill-round buffer overflowed " + f"phase=common_round_lag written=23 processed=7 capacity={monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY} " + "common_round_end=19" + ) + + +def test_gather_round_stats_only_allocates_receive_buffers_on_rank_zero(monkeypatch): + local_round_stats = torch.tensor([[[3, 128]]], dtype=torch.int64) + sentinel_group = object() + + rank_zero_monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + rank_zero_monitor.global_rank = 0 + rank_zero_monitor.world_size = 2 + rank_zero_monitor.gloo_group = sentinel_group + + def root_gather(input_tensor, gather_list, dst, group): + assert dst == 0 and group is sentinel_group + assert len(gather_list) == 2 + gather_list[0].copy_(input_tensor) + gather_list[1].copy_(input_tensor + 1) + + monkeypatch.setattr(monitor_module.dist, "gather", root_gather) + result = rank_zero_monitor._gather_round_stats(local_round_stats) + assert torch.equal(result, torch.tensor([[[[3, 128], [4, 129]]]], dtype=torch.int64)) + + nonzero_monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + nonzero_monitor.global_rank = 1 + nonzero_monitor.world_size = 2 + nonzero_monitor.gloo_group = sentinel_group + + def nonroot_gather(input_tensor, gather_list, dst, group): + assert input_tensor is local_round_stats + assert gather_list is None + assert dst == 0 and group is sentinel_group + + monkeypatch.setattr(monitor_module.dist, "gather", nonroot_gather) + assert nonzero_monitor._gather_round_stats(local_round_stats) is None + + +def test_find_fused_moe_weights_discovers_any_layer_member_once_and_sorts(monkeypatch): + class FakeFusedMoeWeight: + def __init__(self, layer_num, enabled=True): + self.layer_num_ = layer_num + self.enable_ep_moe = enabled + + monkeypatch.setattr(monitor_module, "FusedMoeWeight", FakeFusedMoeWeight) + first = FakeFusedMoeWeight(3) + second = FakeFusedMoeWeight(1) + disabled = FakeFusedMoeWeight(0, enabled=False) + model = SimpleNamespace( + trans_layers_weight=[ + SimpleNamespace(experts_=first, alias=first, ignored=disabled), + SimpleNamespace(any_direct_member=second), + ] + ) + + assert monitor_module._find_fused_moe_weights(model) == [second, first] + + +def test_monitor_disable_detaches_counters_from_all_impls(): + impl = SimpleNamespace(ep_balance_counters="unset") + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.weights = [SimpleNamespace(fuse_moe_impl=impl)] + monitor.enabled = True + monitor._disable() + assert impl.ep_balance_counters is None + assert not monitor.enabled + + +def test_critical_overhead_requires_minimum_samples_for_every_layer(): + round_stats = torch.tensor([[[[200, 32], [200, 32]], [[1, 32], [1, 32]]]], dtype=torch.int64) + assert ( + calculate_prefill_balance_stats( + round_stats, + layer_routed_experts=torch.tensor([1, 1]), + layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), + layer_topks=torch.tensor([2.0, 2.0]), + source_token_replication=1, + ) + is None + ) + + +def test_ep_moe_normal_and_prefill_enable_monitor_by_default(): + assert should_enable_ep_balance_monitor(_monitor_args()) + assert should_enable_ep_balance_monitor(_monitor_args(run_mode="prefill")) + + +def test_disable_ep_balance_monitor_turns_monitor_off(): + assert not should_enable_ep_balance_monitor(_monitor_args(disable_ep_balance_monitor=True)) + + +def test_non_ep_moe_and_decode_mode_do_not_enable_monitor(): + assert not should_enable_ep_balance_monitor(_monitor_args(enable_ep_moe=False)) + assert not should_enable_ep_balance_monitor(_monitor_args(run_mode="decode")) + + +def test_prefill_cudagraph_silently_disables_monitor(): + assert not should_enable_ep_balance_monitor(_monitor_args(run_mode="prefill", enable_prefill_cudagraph=True)) + + +def test_sm100_silently_disables_monitor(monkeypatch): + monkeypatch.setattr(monitor_module, "is_sm100_gpu", lambda: True) + assert not should_enable_ep_balance_monitor(_monitor_args()) From 5a0224683556efd48c560aa3c44110d3f59bbb6f Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 5 Aug 2026 17:53:26 +0800 Subject: [PATCH 02/72] feat: remove auto_update_redundancy_expert function --- docs/CN/source/tutorial/api_server_args.rst | 11 - docs/EN/source/tutorial/api_server_args.rst | 11 - .../meta_weights/fused_moe/ep_redundancy.py | 195 ------------------ .../fused_moe/fused_moe_weight.py | 11 +- .../meta_weights/fused_moe/impl/base_impl.py | 4 - .../fused_moe/impl/deepgemm_impl.py | 12 -- .../fused_moe/impl/triton_impl.py | 4 - .../redundancy_topk_ids_repair.py | 111 ---------- lightllm/distributed/communication_op.py | 3 +- lightllm/server/api_cli.py | 11 - lightllm/server/core/objs/start_args_type.py | 2 - .../mode_backend/redundancy_expert_manager.py | 158 -------------- .../server/router/model_infer/model_rpc.py | 8 - lightllm/utils/envs_utils.py | 60 ------ .../test_redundancy_expert_config.json | 180 ---------------- .../test_redundancy_topk_ids_repair.py | 151 -------------- 16 files changed, 3 insertions(+), 929 deletions(-) delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py delete mode 100644 lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py delete mode 100644 lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py delete mode 100644 test/advanced_config/redundancy_expert/test_redundancy_expert_config.json delete mode 100644 unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index eee60082ca..53027e97a6 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -732,17 +732,6 @@ MTP 多预测参数 增加此值允许更多预测,但确保模型与指定的步数兼容。 目前 deepseekv3/r1 模型仅支持 1 步 -DeepSeek 冗余专家参数 ---------------------- - -.. option:: --ep_redundancy_expert_config_path - - 冗余专家配置的路径。可用于 deepseekv3 模型。 - -.. option:: --auto_update_redundancy_expert - - 是否通过在线专家使用计数器为 deepseekv3 模型更新冗余专家。 - 监控和日志参数 -------------- diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index e4c9151c0d..0f82294c61 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -748,17 +748,6 @@ MTP Multi-Prediction Parameters Increasing this value allows more predictions, but ensure the model is compatible with the specified number of steps. Currently deepseekv3/r1 models only support 1 step -DeepSeek Redundant Expert Parameters ------------------------------------- - -.. option:: --ep_redundancy_expert_config_path - - Path to redundant expert configuration. Can be used for deepseekv3 models. - -.. option:: --auto_update_redundancy_expert - - Whether to update redundant experts for deepseekv3 models through online expert usage counters. - Monitoring and Logging Parameters --------------------------------- diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py deleted file mode 100644 index 749400c8d8..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py +++ /dev/null @@ -1,195 +0,0 @@ -import numpy as np -import torch -from .fused_moe_weight import FusedMoeWeight -from lightllm.utils.log_utils import init_logger -from typing import Dict - -logger = init_logger(__name__) - - -class FusedMoeWeightEPAutoRedundancy: - def __init__( - self, - ep_fused_moe_weight: FusedMoeWeight, - ) -> None: - super().__init__() - self._ep_w = ep_fused_moe_weight - self.redundancy_expert_num = self._ep_w.redundancy_expert_num - - def clear_counter(self): - self._ep_w.routed_expert_counter_tensor.fill_(0) - return - - def prepare_redundancy_experts( - self, - ): - expert_counter = self._ep_w.routed_expert_counter_tensor.detach().cpu().numpy() - logger.info( - f"layer_index {self._ep_w.layer_num_} global_rank {self._ep_w.global_rank_}" - f" expert_counter: {expert_counter}" - ) - self._ep_w.routed_expert_counter_tensor.fill_(0) - ep_n_routed_experts = self._ep_w.n_routed_experts // self._ep_w.global_world_size - start_expert_id = ep_n_routed_experts * self._ep_w.global_rank_ - no_redundancy_expert_ids = list(range(start_expert_id, start_expert_id + ep_n_routed_experts)) - - # 统计 0 rank 上的全局 topk 冗余信息,帮助导出一份全局可用的静态使用的冗余专家静态配置。 - if self._ep_w.global_rank_ == 0: - # int(e) for serialization, int64 can not be serialized by json.dump. - topk_redundancy_expert_ids = list(int(e) for e in np.argsort(expert_counter)[-self.redundancy_expert_num :]) - else: - topk_redundancy_expert_ids = None - - # 不要选中当前已经存在的非冗余专家作为冗余专家 - expert_counter[no_redundancy_expert_ids] = 0 - - self.redundancy_expert_ids = list(np.argsort(expert_counter)[-self.redundancy_expert_num :]) - logger.info( - f"layer_index {self._ep_w.layer_num_} global_rank {self._ep_w.global_rank_}" - f" new select redundancy_expert_ids : {self.redundancy_expert_ids}" - ) - - # 准备加载过度变量。 - self.experts_up_projs = [None] * self.redundancy_expert_num - self.experts_gate_projs = [None] * self.redundancy_expert_num - self.experts_up_proj_scales = [None] * self.redundancy_expert_num - self.experts_gate_proj_scales = [None] * self.redundancy_expert_num - self.w2_list = [None] * self.redundancy_expert_num - self.w2_scale_list = [None] * self.redundancy_expert_num - self.w13 = [None, None] # weight, weight_scale - self.w2 = [None, None] # weight, weight_scale - return topk_redundancy_expert_ids - - def load_hf_weights(self, weights): - # 加载冗余专家的权重参数 - for i, redundant_expert_id in enumerate(self.redundancy_expert_ids): - i_experts = redundant_expert_id - w1_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w1_weight_name}.weight" - w2_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w2_weight_name}.weight" - w3_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w3_weight_name}.weight" - if w1_weight in weights: - self.experts_gate_projs[i] = weights[w1_weight] - if w3_weight in weights: - self.experts_up_projs[i] = weights[w3_weight] - if w2_weight in weights: - self.w2_list[i] = weights[w2_weight] - - self._load_weight_scale(weights) - self._fuse() - - def _fuse(self): - self._fuse_weight_scale() - - with self._ep_w.lock: - if ( - hasattr(self, "experts_up_projs") - and None not in self.experts_up_projs - and None not in self.experts_gate_projs - and None not in self.w2_list - ): - gate_out_dim, gate_in_dim = self.experts_gate_projs[0].shape - up_out_dim, up_in_dim = self.experts_up_projs[0].shape - assert gate_in_dim == up_in_dim - dtype = self.experts_gate_projs[0].dtype - total_expert_num = self.redundancy_expert_num - - w13 = torch.empty((total_expert_num, gate_out_dim + up_out_dim, gate_in_dim), dtype=dtype, device="cpu") - - for i_experts in range(self.redundancy_expert_num): - w13[i_experts, 0:gate_out_dim:, :] = self.experts_gate_projs[i_experts] - w13[i_experts, gate_out_dim:, :] = self.experts_up_projs[i_experts] - - inter_shape, hidden_size = self.w2_list[0].shape[0], self.w2_list[0].shape[1] - w2 = torch._utils._flatten_dense_tensors(self.w2_list).view(len(self.w2_list), inter_shape, hidden_size) - if self._ep_w.quant_method._check_weight_need_quanted(weight=w13): - w13_pack, _ = self._ep_w.quant_method.create_moe_weight( - out_dims=[gate_out_dim + up_out_dim], - in_dim=1, - dtype=self._ep_w.data_type_, - device_id=self._ep_w.device_id_, - num_experts=self.redundancy_expert_num, - ) - self._ep_w.quant_method.quantize(w13, w13_pack) - w2_pack, _ = self._ep_w.quant_method.create_moe_weight( - out_dims=[inter_shape], - in_dim=hidden_size, - dtype=self._ep_w.data_type_, - device_id=self._ep_w.device_id_, - num_experts=self.redundancy_expert_num, - ) - self._ep_w.quant_method.quantize(w2, w2_pack) - - self.w13[0] = w13_pack.weight - self.w13[1] = w13_pack.weight_scale - self.w2[0] = w2_pack.weight - self.w2[1] = w2_pack.weight_scale - else: - self.w13[0] = w13 - self.w2[0] = w2 - delattr(self, "w2_list") - delattr(self, "experts_up_projs") - delattr(self, "experts_gate_projs") - - def _fuse_weight_scale(self): - with self._ep_w.lock: - if ( - hasattr(self, "experts_up_proj_scales") - and None not in self.experts_up_proj_scales - and None not in self.experts_gate_proj_scales - and None not in self.w2_scale_list - ): - gate_out_dim, gate_in_dim = self.experts_gate_proj_scales[0].shape - up_out_dim, up_in_dim = self.experts_up_proj_scales[0].shape - assert gate_in_dim == up_in_dim - dtype = self.experts_gate_proj_scales[0].dtype - total_expert_num = self.redundancy_expert_num - w13_scale = torch.empty( - (total_expert_num, gate_out_dim + up_out_dim, gate_in_dim), dtype=dtype, device="cpu" - ) - for i_experts in range(self.redundancy_expert_num): - w13_scale[i_experts, 0:gate_out_dim:, :] = self.experts_gate_proj_scales[i_experts] - w13_scale[i_experts, gate_out_dim:, :] = self.experts_up_proj_scales[i_experts] - - inter_shape, hidden_size = self.w2_scale_list[0].shape[0], self.w2_scale_list[0].shape[1] - w2_scale = torch._utils._flatten_dense_tensors(self.w2_scale_list).view( - len(self.w2_scale_list), inter_shape, hidden_size - ) - self.w13[1] = w13_scale - self.w2[1] = w2_scale - delattr(self, "w2_scale_list") - delattr(self, "experts_up_proj_scales") - delattr(self, "experts_gate_proj_scales") - - def _load_weight_scale(self, weights: Dict[str, torch.Tensor]) -> None: - # 加载冗余专家的scale参数 - for i, redundant_expert_id in enumerate(self.redundancy_expert_ids): - i_experts = redundant_expert_id - weight_scale_suffix = self._ep_w.quant_method.weight_scale_suffix - w1_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w1_weight_name}.{weight_scale_suffix}" - w2_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w2_weight_name}.{weight_scale_suffix}" - w3_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w3_weight_name}.{weight_scale_suffix}" - if w1_scale in weights: - self.experts_gate_proj_scales[i] = weights[w1_scale] - if w3_scale in weights: - self.experts_up_proj_scales[i] = weights[w3_scale] - if w2_scale in weights: - self.w2_scale_list[i] = weights[w2_scale] - - def commit(self): - for index, dest_tensor in enumerate([self._ep_w.w13.weight, self._ep_w.w13.weight_scale]): - if dest_tensor is not None: - assert isinstance( - dest_tensor, torch.Tensor - ), f"dest_tensor should be a torch.Tensor, but got {type(dest_tensor)}" - dest_tensor[-self.redundancy_expert_num :, :, :] = self.w13[index][:, :, :] - - for index, dest_tensor in enumerate([self._ep_w.w2.weight, self._ep_w.w2.weight_scale]): - if dest_tensor is not None: - assert isinstance( - dest_tensor, torch.Tensor - ), f"dest_tensor should be a torch.Tensor, but got {type(dest_tensor)}" - dest_tensor[-self.redundancy_expert_num :, :, :] = self.w2[index][:, :, :] - - self._ep_w.redundancy_expert_ids_tensor.copy_( - torch.tensor(self.redundancy_expert_ids, dtype=torch.int64, device="cpu") - ) 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 7f369c4fd8..6fc81dd6f9 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 @@ -11,7 +11,6 @@ from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import select_fuse_moe_impl from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback from lightllm.common.quantization.quantize_method import QuantizationMethod -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 @@ -64,9 +63,7 @@ def __init__( routed_scaling_factor=self.routed_scaling_factor, quant_method=self.quant_method, redundancy_expert_num=self.redundancy_expert_num, - redundancy_expert_ids_tensor=self.redundancy_expert_ids_tensor, routed_expert_counter_tensor=self.routed_expert_counter_tensor, - auto_update_redundancy_expert=self.auto_update_redundancy_expert, ) self.lock = threading.Lock() self._create_weight() @@ -81,13 +78,9 @@ def _init_config(self, network_config: Dict[str, Any]): self.scoring_func = network_config.get("scoring_func", "softmax") def _init_redundancy_expert_params(self): - self.redundancy_expert_num = get_redundancy_expert_num() - self.redundancy_expert_ids = get_redundancy_expert_ids(self.layer_num_) - self.auto_update_redundancy_expert: bool = get_env_start_args().auto_update_redundancy_expert - self.redundancy_expert_ids_tensor = torch.tensor(self.redundancy_expert_ids, dtype=torch.int64, device="cuda") + self.redundancy_expert_num = 0 + self.redundancy_expert_ids = [] self.routed_expert_counter_tensor = torch.zeros((self.n_routed_experts,), dtype=torch.int64, device="cuda") - # TODO: find out the reason of failure of deepep when redundancy_expert_num is 1. - assert self.redundancy_expert_num != 1, "redundancy_expert_num can not be 1 for some unknown hang of deepep." def _init_parallel_params(self): if self.enable_ep_moe: diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py index 1e3ad4b196..9b6c42af79 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py @@ -19,9 +19,7 @@ def __init__( routed_scaling_factor: float, quant_method: QuantizationMethod, redundancy_expert_num: int, - redundancy_expert_ids_tensor: torch.Tensor, routed_expert_counter_tensor: torch.Tensor, - auto_update_redundancy_expert: bool, ): self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts @@ -36,9 +34,7 @@ def __init__( # redundancy expert related self.redundancy_expert_num = redundancy_expert_num - self.redundancy_expert_ids_tensor = redundancy_expert_ids_tensor self.routed_expert_counter_tensor = routed_expert_counter_tensor - self.auto_update_redundancy_expert = auto_update_redundancy_expert # workspace for kernel optimization self.workspace = self.create_workspace() 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 ec2c1fd39c..6fabd3304c 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 @@ -16,7 +16,6 @@ ) 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, AutotuneKernelType -from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair class FuseMoeDeepGEMM(FuseMoeTriton): @@ -58,17 +57,6 @@ def _select_experts( if per_expert_scale is not None: topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype) origin_topk_ids = topk_ids - if self.redundancy_expert_num > 0: - # 因为 redundancy_topk_ids_repair 会修改 topk_ids,所以需要先复制一份 - origin_topk_ids = topk_ids.clone() - redundancy_topk_ids_repair( - topk_ids=topk_ids, - redundancy_expert_ids=self.redundancy_expert_ids_tensor, - ep_expert_num=self.ep_n_routed_experts, - global_rank=self.global_rank_, - expert_counter=self.routed_expert_counter_tensor, - enable_counter=self.auto_update_redundancy_expert, - ) return topk_weights, topk_ids, origin_topk_ids def _fused_experts( diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index 1d6a38c069..abe0112004 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -13,9 +13,7 @@ def __init__( routed_scaling_factor: float, quant_method: QuantizationMethod, redundancy_expert_num: int, - redundancy_expert_ids_tensor: torch.Tensor, routed_expert_counter_tensor: torch.Tensor, - auto_update_redundancy_expert: bool, ): super().__init__( n_routed_experts=n_routed_experts, @@ -23,9 +21,7 @@ def __init__( routed_scaling_factor=routed_scaling_factor, quant_method=quant_method, redundancy_expert_num=redundancy_expert_num, - redundancy_expert_ids_tensor=redundancy_expert_ids_tensor, routed_expert_counter_tensor=routed_expert_counter_tensor, - auto_update_redundancy_expert=auto_update_redundancy_expert, ) def create_workspace(self): diff --git a/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py b/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py deleted file mode 100644 index ba48f414db..0000000000 --- a/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py +++ /dev/null @@ -1,111 +0,0 @@ -import torch -import triton -import triton.language as tl - - -@triton.jit -def _redundancy_topk_ids_repair_kernel( - topk_ids_ptr, - topk_total_num, - ep_expert_num, - redundancy_expert_num, - global_rank, - redundancy_expert_ids_ptr, - expert_counter_ptr, - BLOCK_SIZE: tl.constexpr, - ENABLE_COUNTER: tl.constexpr, -): - block_index = tl.program_id(0) - offs_d = block_index * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offs_d < topk_total_num - current_topk_ids = tl.load(topk_ids_ptr + offs_d, mask=mask, other=0) - - if ENABLE_COUNTER: - tl.atomic_add(expert_counter_ptr + current_topk_ids, 1, mask=mask) - - # Remap original expert IDs to a new space that accounts for redundant expert slots. - new_current_topk_ids = (current_topk_ids // ep_expert_num) * redundancy_expert_num + current_topk_ids - - for i in tl.range(0, redundancy_expert_num, step=1, num_stages=3): - cur_redundancy_expert_id = tl.load(redundancy_expert_ids_ptr + i) - cur_redundancy_expert_id = ( - cur_redundancy_expert_id // ep_expert_num - ) * redundancy_expert_num + cur_redundancy_expert_id - new_current_topk_ids = tl.where( - new_current_topk_ids == cur_redundancy_expert_id, - (ep_expert_num + redundancy_expert_num) * (global_rank) + ep_expert_num + i, - new_current_topk_ids, - ) - - tl.store(topk_ids_ptr + offs_d, new_current_topk_ids, mask=mask) - return - - -@torch.no_grad() -def redundancy_topk_ids_repair( - topk_ids: torch.Tensor, - redundancy_expert_ids: torch.Tensor, - ep_expert_num: int, - global_rank: int, - expert_counter: torch.Tensor = None, - enable_counter: bool = False, -): - assert topk_ids.is_contiguous() - assert len(topk_ids.shape) == 2 - assert redundancy_expert_ids is not None - redundancy_expert_num = redundancy_expert_ids.shape[0] - BLOCK_SIZE = 512 - grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),) - num_warps = 4 - - _redundancy_topk_ids_repair_kernel[grid]( - topk_ids_ptr=topk_ids, - topk_total_num=topk_ids.numel(), - ep_expert_num=ep_expert_num, - redundancy_expert_num=redundancy_expert_num, - global_rank=global_rank, - redundancy_expert_ids_ptr=redundancy_expert_ids, - expert_counter_ptr=expert_counter, - BLOCK_SIZE=BLOCK_SIZE, - ENABLE_COUNTER=enable_counter, - num_warps=num_warps, - num_stages=3, - ) - return - - -@triton.jit -def _expert_id_counter_kernel( - topk_ids_ptr, - topk_total_num, - expert_counter_ptr, - BLOCK_SIZE: tl.constexpr, -): - block_index = tl.program_id(0) - offs_d = block_index * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offs_d < topk_total_num - current_topk_ids = tl.load(topk_ids_ptr + offs_d, mask=mask, other=0) - tl.atomic_add(expert_counter_ptr + current_topk_ids, 1, mask=mask) - return - - -@torch.no_grad() -def expert_id_counter( - topk_ids: torch.Tensor, - expert_counter: torch.Tensor, -): - assert topk_ids.is_contiguous() - assert len(topk_ids.shape) == 2 - BLOCK_SIZE = 512 - grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),) - num_warps = 4 - - _expert_id_counter_kernel[grid]( - topk_ids_ptr=topk_ids, - topk_total_num=topk_ids.numel(), - expert_counter_ptr=expert_counter, - BLOCK_SIZE=BLOCK_SIZE, - num_warps=num_warps, - num_stages=1, - ) - return diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 76b202f4d5..b721585596 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -30,7 +30,6 @@ get_env_start_args, get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, - get_redundancy_expert_num, ) from lightllm.utils.dist_utils import ( get_global_world_size, @@ -202,7 +201,7 @@ def new_deepep_group( self.ll_num_tokens = prefill_num_max_dispatch_tokens_per_rank self.ll_decode_num_tokens = decode_num_max_dispatch_tokens_per_rank self.ll_hidden = hidden_size - self.ll_num_experts = n_routed_experts + get_redundancy_expert_num() * global_world_size + self.ll_num_experts = n_routed_experts self.ep_buffer = deep_ep.ElasticBuffer( deepep_group, num_max_tokens_per_rank=self.ll_num_tokens, diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 93656f600b..9d90ff69ad 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -771,17 +771,6 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="""Disable the prefill expert balance monitor enabled by default for EP-MoE.""", ) - parser.add_argument( - "--ep_redundancy_expert_config_path", - type=str, - default=None, - help="""Path of the redundant expert config. It can be used for deepseekv3 model.""", - ) - parser.add_argument( - "--auto_update_redundancy_expert", - action="store_true", - help="""Whether to update the redundant expert for deepseekv3 model by online expert used counter.""", - ) parser.add_argument( "--enable_fused_shared_experts", action="store_true", diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 16ee99b3c7..6ad7a9b9d8 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -187,8 +187,6 @@ class StartArgs: ) enable_ep_moe: bool = field(default=False) disable_ep_balance_monitor: bool = field(default=False) - ep_redundancy_expert_config_path: Optional[str] = field(default=None) - auto_update_redundancy_expert: bool = field(default=False) enable_fused_shared_experts: bool = field(default=False) mtp_mode: Optional[str] = field( default=None, diff --git a/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py b/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py deleted file mode 100644 index 596eca4f24..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py +++ /dev/null @@ -1,158 +0,0 @@ -# 对于 deepseekv3 模型在 ep 运行模式下,自动分析统计各个专家的出现频率,然后 -# 自动更新当前的冗余专家为新的冗余专家。 -import torch -import time -import enum -import lightllm.utils.petrel_helper as utils -import threading -import json -from typing import List -from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_redundancy import ( - FusedMoeWeightEPAutoRedundancy, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight -from lightllm.utils.envs_utils import get_env_start_args, get_redundancy_expert_update_interval -from lightllm.utils.envs_utils import get_redundancy_expert_update_max_load_count -from lightllm.utils.envs_utils import get_redundancy_expert_num -from lightllm.utils.dist_utils import get_global_rank -from lightllm.common.basemodel.layer_weights.hf_load_utils import load_func -from lightllm.utils.log_utils import init_logger - -logger = init_logger(__name__) - - -class RedundancyExpertManager: - def __init__(self, model: TpPartBaseModel): - self.args = get_env_start_args() - self.model = model - self.ep_fused_moeweights: List[FusedMoeWeightEPAutoRedundancy] = [] - for layer in self.model.trans_layers_weight: - ep_weights = self._find_members_of_class(layer, FusedMoeWeight) - assert len(ep_weights) <= 1 - self.ep_fused_moeweights.extend([FusedMoeWeightEPAutoRedundancy(e) for e in ep_weights]) - - # save load params - self.use_safetensors = True - files = utils.PetrelHelper.list(self.args.model_dir, extension="all") - candidate_files = list(filter(lambda x: x.endswith(".safetensors"), files)) - if len(candidate_files) == 0: - self.use_safetensors = False - candidate_files = list(filter(lambda x: x.endswith(".bin"), files)) - assert len(candidate_files) != 0, "can only support pytorch tensor and safetensors format for weights." - self.candidate_files = candidate_files - - # state 1. check_to_update 2. prepare_update 3. start_load_hf_weights 4. wait_load_ready, 5. commit - self.state: _STATE = _STATE.CHECK_TO_UPDATE - self.update_time = time.time() - self.update_interval = get_redundancy_expert_update_interval() - self.load_thread: threading.Thread = None - self.global_rank = get_global_rank() - # 冗余专家的最大加载次数 - self.load_count = 0 - self.max_load_count = get_redundancy_expert_update_max_load_count() - - # 清理counter - self._clear_all_counter() - - self.rank0_redundancy_expert_config = { - "redundancy_expert_num": get_redundancy_expert_num(), - "default": list(range(get_redundancy_expert_num())), - } - - def step(self): - if self.load_count >= self.max_load_count: - return - - if self.state == _STATE.CHECK_TO_UPDATE: - cur_time = time.time() - if cur_time - self.update_time > self.update_interval: - self.update_time = cur_time - self.state = _STATE.PREPARE_UPDATE - logger.info(f"global_rank {self.global_rank} state to prepare update") - elif self.state == _STATE.PREPARE_UPDATE: - self._prepare_load_new_redundancy_expert() - self.state = _STATE.START_LOAD_HF_WEIGHTS - logger.info(f"global_rank {self.global_rank} state to start load hf weights") - - elif self.state == _STATE.START_LOAD_HF_WEIGHTS: - self.load_thread = threading.Thread(target=self._load_hf_weights, daemon=True) - self.load_thread.start() - self.state = _STATE.WAIT_LOAD_READY - logger.info(f"global_rank {self.global_rank} state to wait load ready") - - elif self.state == _STATE.WAIT_LOAD_READY: - if not self.load_thread.is_alive(): - self.load_thread = None - self.state = _STATE.COMMIT - logger.info(f"global_rank {self.global_rank} state to commit") - - elif self.state == _STATE.COMMIT: - self._commit() - self.state = _STATE.CHECK_TO_UPDATE - self.load_count += 1 - logger.info(f"global_rank {self.global_rank} state to check to update") - return - - def _prepare_load_new_redundancy_expert(self): - for w in self.ep_fused_moeweights: - topk_redundancy_expert_ids = w.prepare_redundancy_experts() - if self.global_rank == 0: - self.rank0_redundancy_expert_config[str(w._ep_w.layer_num)] = topk_redundancy_expert_ids - - if self.global_rank == 0: - try: - with open("./redundancy_expert_config.json", "w") as f: - json.dump(self.rank0_redundancy_expert_config, f, indent=4) - logger.info( - f"rank {self.global_rank} save redundancy_expert_config.json to ./redundancy_expert_config.json" - ) - except BaseException as e: - logger.exception(str(e)) - logger.error(f"global rank {self.global_rank} save redundancy_expert_config.json failed") - - return - - def _load_hf_weights(self): - start = time.time() - try: - for file in self.candidate_files: - load_func( - file, - use_safetensors=self.use_safetensors, - pre_post_layer=None, - transformer_layer_list=self.ep_fused_moeweights, - weight_dir=self.args.model_dir, - ) - except BaseException as e: - logger.exception(str(e)) - raise e - cost_time = time.time() - start - logger.info(f"global rank {self.global_rank} load redundancy_expert cost time: {cost_time} s") - return - - def _commit(self): - for w in self.ep_fused_moeweights: - w.commit() - return - - def _find_members_of_class(self, obj, cls): - members = [] - for attr in dir(obj): - value = getattr(obj, attr) - if isinstance(value, cls): - members.append(value) - return members - - def _clear_all_counter(self): - for w in self.ep_fused_moeweights: - w.clear_counter() - return - - -class _STATE(enum.Enum): - CHECK_TO_UPDATE = 0 - PREPARE_UPDATE = 1 - START_LOAD_HF_WEIGHTS = 2 - WAIT_LOAD_READY = 3 - COMMIT = 4 diff --git a/lightllm/server/router/model_infer/model_rpc.py b/lightllm/server/router/model_infer/model_rpc.py index 3ae4d4cbc2..e4784e79bd 100644 --- a/lightllm/server/router/model_infer/model_rpc.py +++ b/lightllm/server/router/model_infer/model_rpc.py @@ -25,7 +25,6 @@ PDDecodeNode, PDDPForDecodeNode, ) -from lightllm.server.router.model_infer.mode_backend.redundancy_expert_manager import RedundancyExpertManager from lightllm.server.router.model_infer.mode_backend.rl_backend_ops import RlBackendOps from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( EPBalanceMonitor, @@ -102,13 +101,6 @@ def exposed_init_model(self, kvargs): self.backend.init_model(kvargs) self.rl_backend_ops = RlBackendOps(self.backend) if self.args.enable_rl else None - # only deepseekv3 can support auto_update_redundancy_expert - if self.args.auto_update_redundancy_expert: - self.redundancy_expert_manager = RedundancyExpertManager(self.backend.model) - logger.info("init redundancy_expert_manager") - else: - self.redundancy_expert_manager = None - if should_enable_ep_balance_monitor(self.args): monitor = EPBalanceMonitor(self.backend.model) if monitor.enabled: diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index acd2ac6711..cc7e968193 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -96,66 +96,6 @@ def get_lightllm_websocket_max_message_size(): return int(os.getenv("LIGHTLLM_WEBSOCKET_MAX_SIZE", 128 * 1024 * 1024)) -# get_redundancy_expert_ids and get_redundancy_expert_num are primarily -# used to obtain the IDs and number of redundant experts during inference. -# They depend on a configuration file specified by ep_redundancy_expert_config_path, -# which is a JSON formatted text file. -# The content format is as follows: -# { -# "redundancy_expert_num": 1, # Number of redundant experts per rank -# "0": [0], # Key: layer_index (string), -# # Value: list of original expert IDs that are redundant for this layer -# "1": [0], -# "default": [0] # Default list of redundant expert IDs if layer-specific entry is not found -# } - - -@lru_cache(maxsize=None) -def get_redundancy_expert_ids(layer_index: int): - """ - Get the redundancy expert ids from the environment variable. - :return: List of redundancy expert ids. - """ - args = get_env_start_args() - if args.ep_redundancy_expert_config_path is None: - return [] - - with open(args.ep_redundancy_expert_config_path, "r") as f: - config = json.load(f) - if str(layer_index) in config: - return config[str(layer_index)] - else: - return config.get("default", []) - - -@lru_cache(maxsize=None) -def get_redundancy_expert_num(): - """ - Get the number of redundancy experts from the environment variable. - :return: Number of redundancy experts. - """ - args = get_env_start_args() - if args.ep_redundancy_expert_config_path is None: - return 0 - - with open(args.ep_redundancy_expert_config_path, "r") as f: - config = json.load(f) - if "redundancy_expert_num" in config: - return config["redundancy_expert_num"] - else: - return 0 - - -@lru_cache(maxsize=None) -def get_redundancy_expert_update_interval(): - return int(os.getenv("LIGHTLLM_REDUNDANCY_EXPERT_UPDATE_INTERVAL", 30 * 60)) - - -@lru_cache(maxsize=None) -def get_redundancy_expert_update_max_load_count(): - return int(os.getenv("LIGHTLLM_REDUNDANCY_EXPERT_UPDATE_MAX_LOAD_COUNT", 1)) - - @lru_cache(maxsize=None) def get_triton_autotune_level(): return int(os.getenv("LIGHTLLM_TRITON_AUTOTUNE_LEVEL", 0)) diff --git a/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json b/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json deleted file mode 100644 index 241ab25ea3..0000000000 --- a/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json +++ /dev/null @@ -1,180 +0,0 @@ -{ - "redundancy_expert_num": 1, - "default": [ - 0 - ], - "3": [ - 226 - ], - "4": [ - 123 - ], - "5": [ - 187 - ], - "6": [ - 138 - ], - "7": [ - 132 - ], - "8": [ - 240 - ], - "9": [ - 4 - ], - "10": [ - 88 - ], - "11": [ - 60 - ], - "12": [ - 161 - ], - "13": [ - 178 - ], - "14": [ - 80 - ], - "15": [ - 144 - ], - "16": [ - 195 - ], - "17": [ - 251 - ], - "18": [ - 226 - ], - "19": [ - 87 - ], - "20": [ - 149 - ], - "21": [ - 45 - ], - "22": [ - 214 - ], - "23": [ - 41 - ], - "24": [ - 46 - ], - "25": [ - 156 - ], - "26": [ - 112 - ], - "27": [ - 185 - ], - "28": [ - 58 - ], - "29": [ - 156 - ], - "30": [ - 147 - ], - "31": [ - 199 - ], - "32": [ - 16 - ], - "33": [ - 188 - ], - "34": [ - 227 - ], - "35": [ - 136 - ], - "36": [ - 84 - ], - "37": [ - 15 - ], - "38": [ - 204 - ], - "39": [ - 96 - ], - "40": [ - 226 - ], - "41": [ - 25 - ], - "42": [ - 69 - ], - "43": [ - 122 - ], - "44": [ - 152 - ], - "45": [ - 113 - ], - "46": [ - 98 - ], - "47": [ - 68 - ], - "48": [ - 13 - ], - "49": [ - 102 - ], - "50": [ - 214 - ], - "51": [ - 201 - ], - "52": [ - 182 - ], - "53": [ - 235 - ], - "54": [ - 162 - ], - "55": [ - 125 - ], - "56": [ - 62 - ], - "57": [ - 121 - ], - "58": [ - 105 - ], - "59": [ - 236 - ], - "60": [ - 117 - ] -} \ No newline at end of file diff --git a/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py b/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py deleted file mode 100644 index 16131ef935..0000000000 --- a/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py +++ /dev/null @@ -1,151 +0,0 @@ -import torch -import pytest -from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair -from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import expert_id_counter -from lightllm.utils.log_utils import init_logger - -logger = init_logger(__name__) - - -def test_redundancy_topk_ids_repair(): - ep_expert_num = 4 - global_rank = 0 - redundancy_expert_num = 1 - topk_ids = torch.tensor( - [ - [0, 1, 2, 3], - [7, 9, 10, 11], - [1, 3, 5, 7], - ], - dtype=torch.int64, - device="cuda", - ) - - redundancy_expert_ids = torch.tensor( - [ - 0, - ], - dtype=torch.int64, - device="cuda", - ) - - expert_id_counter = torch.zeros(12, dtype=torch.int64, device="cuda") - - redundancy_topk_ids_repair( - topk_ids=topk_ids, - redundancy_expert_ids=redundancy_expert_ids, - ep_expert_num=ep_expert_num, - global_rank=global_rank, - expert_counter=expert_id_counter, - enable_counter=True, - ) - - ans_topk_ids = torch.tensor( - [ - [0, 1, 2, 3], - [7, 9, 10, 11], - [1, 3, 5, 7], - ], - dtype=torch.int64, - device="cuda", - ) - ans_topk_ids = (ans_topk_ids // ep_expert_num) * redundancy_expert_num + ans_topk_ids - new_redundancy_expert_ids = (redundancy_expert_ids // ep_expert_num) * redundancy_expert_num + redundancy_expert_ids - ans_topk_ids[ans_topk_ids == new_redundancy_expert_ids[0]] = ( - (ep_expert_num + redundancy_expert_num) * global_rank + ep_expert_num + 0 - ) - - assert torch.equal(topk_ids, ans_topk_ids) - assert torch.equal( - expert_id_counter, torch.tensor([1, 2, 1, 2, 0, 1, 0, 2, 0, 1, 1, 1], dtype=torch.int64, device="cuda") - ) - - ep_expert_num = 4 - global_rank = 1 - redundancy_expert_num = 1 - topk_ids = torch.tensor( - [ - [0, 1, 2, 3], - [7, 9, 10, 11], - [1, 3, 5, 7], - ], - dtype=torch.int64, - device="cuda", - ) - - redundancy_expert_ids = torch.tensor( - [ - 5, - ], - dtype=torch.int64, - device="cuda", - ) - redundancy_topk_ids_repair( - topk_ids=topk_ids, - redundancy_expert_ids=redundancy_expert_ids, - ep_expert_num=ep_expert_num, - global_rank=global_rank, - ) - - ans_topk_ids = torch.tensor( - [ - [0, 1, 2, 3], - [7, 9, 10, 11], - [1, 3, 5, 7], - ], - dtype=torch.int64, - device="cuda", - ) - ans_topk_ids = (ans_topk_ids // ep_expert_num) * redundancy_expert_num + ans_topk_ids - new_redundancy_expert_ids = (redundancy_expert_ids // ep_expert_num) * redundancy_expert_num + redundancy_expert_ids - ans_topk_ids[ans_topk_ids == new_redundancy_expert_ids[0]] = ( - (ep_expert_num + redundancy_expert_num) * global_rank + ep_expert_num + 0 - ) - - assert torch.equal(topk_ids, ans_topk_ids) - - -def test_expert_id_counter(): - token_num = 256 - tok_ids = torch.randint( - low=0, - high=12, - size=(token_num, 8), - dtype=torch.int64, - device="cuda", - ) - expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") - expert_id_counter(topk_ids=tok_ids, expert_counter=expert_counter) - - ans_expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") - ids, counts = torch.unique(tok_ids.view(-1), return_counts=True) - ans_expert_counter[ids] = counts - - assert torch.equal(expert_counter, ans_expert_counter) - - # test speed - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - for _ in range(100): - tok_ids = torch.randint( - low=0, - high=12, - size=(token_num, 8), - dtype=torch.int64, - device="cuda", - ) - expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") - expert_id_counter(topk_ids=tok_ids, expert_counter=expert_counter) - graph.replay() - - start_event = torch.cuda.Event(enable_timing=True) - start_event.record() - graph.replay() - end_event = torch.cuda.Event(enable_timing=True) - end_event.record() - torch.cuda.synchronize() - logger.info(f"expert_id_counter time cost: {start_event.elapsed_time(end_event)} ms") - - -if __name__ == "__main__": - pytest.main() From efeede67c94be45927e2f1a3d7155adbe1d2b8d3 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 22 Jul 2026 15:49:26 +0800 Subject: [PATCH 03/72] feat: add EPLB --- lightllm/common/basemodel/basemodel.py | 9 +- .../meta_weights/fused_moe/eplb_placement.py | 570 +++ .../fused_moe/expert_parallel_state.py | 38 + .../fused_moe/fused_moe_weight.py | 97 +- .../meta_weights/fused_moe/impl/__init__.py | 35 +- .../meta_weights/fused_moe/impl/base_impl.py | 87 +- .../fused_moe/impl/deepgemm_impl.py | 114 +- .../fused_moe/impl/marlin_impl.py | 4 + .../fused_moe/impl/triton_impl.py | 73 +- .../triton_kernel/fused_moe/eplb_kernels.py | 55 + .../triton_kernel/fused_moe/grouped_topk.py | 246 ++ .../triton_kernel/fused_moe/topk_select.py | 9 - lightllm/distributed/communication_op.py | 38 +- lightllm/server/api_cli.py | 12 + lightllm/server/api_start.py | 11 + lightllm/server/core/objs/start_args_type.py | 2 + .../model_infer/mode_backend/base_backend.py | 10 +- .../mode_backend/chunked_prefill/impl.py | 5 + .../mode_backend/dp_backend/impl.py | 5 + .../model_infer/mode_backend/eplb_manager.py | 558 +++ .../model_infer/mode_backend/eplb_transfer.py | 655 ++++ lightllm/utils/envs_utils.py | 30 + unit_tests/common/fused_moe/test_eplb.py | 3360 +++++++++++++++++ .../fused_moe/test_eplb_transfer_gpu.py | 225 ++ unit_tests/server/test_api_start_eplb.py | 28 + 25 files changed, 6102 insertions(+), 174 deletions(-) create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py create mode 100644 lightllm/common/basemodel/triton_kernel/fused_moe/eplb_kernels.py create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb_manager.py create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb_transfer.py create mode 100644 unit_tests/common/fused_moe/test_eplb.py create mode 100644 unit_tests/common/fused_moe/test_eplb_transfer_gpu.py create mode 100644 unit_tests/server/test_api_start_eplb.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index b5d364e61a..ea9d95927c 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -70,6 +70,7 @@ class TpPartBaseModel: def __init__(self, kvargs): self.args = get_env_start_args() + self.eplb_manager = None self.ep_balance_monitor = None self.run_mode = kvargs["run_mode"] self.weight_dir_ = kvargs["weight_dir"] @@ -318,13 +319,15 @@ def forward(self, model_input: ModelInput): if model_input.is_prefill: model_output = self._prefill(model_input=model_input) - self._record_prefill_ep_balance() + self._after_prefill() return model_output return self._decode(model_input) - def _record_prefill_ep_balance(self): + def _after_prefill(self): if self.ep_balance_monitor is not None: self.ep_balance_monitor.record_prefill_round() + if self.eplb_manager is not None: + self.eplb_manager.step() def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() @@ -820,7 +823,7 @@ def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input dist_group_manager.clear_deepep_buffer() model_output0.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event model_output1.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event - self._record_prefill_ep_balance() + self._after_prefill() return model_output0, model_output1 @torch.no_grad() diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py new file mode 100644 index 0000000000..a7504e975e --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -0,0 +1,570 @@ +from dataclasses import dataclass +from functools import lru_cache +from typing import Dict, Tuple +import torch + + +def build_initial_redundant_expert_ids( + num_logical_experts: int, + num_ranks: int, + num_redundant_experts_per_rank: int, +) -> torch.Tensor: + """Build a deterministic initial placement without local duplicates.""" + assert num_logical_experts % num_ranks == 0 + num_experts_per_rank = num_logical_experts // num_ranks + assert 0 < num_redundant_experts_per_rank <= num_logical_experts - num_experts_per_rank + + # 初始化结果确定,不依赖随机数。 + # 每个 rank 不会复制自己原本拥有的 expert。 + # 同一个 rank 的冗余槽位不会重复。 + # 最后一个 rank 通过取模自然回绕。 + rank_offsets = torch.arange(1, num_ranks + 1, dtype=torch.int64)[:, None] * num_experts_per_rank + expert_offsets = torch.arange(num_redundant_experts_per_rank, dtype=torch.int64) + return (rank_offsets + expert_offsets) % num_logical_experts + + +def build_logical_to_physical_map( + redundant_expert_ids: torch.Tensor, # 冗余布局,shape 为 [num_ranks, num_redundant_experts_per_rank]。 + num_logical_experts: int, # 逻辑 expert 的总数。 + source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 + node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 +) -> Tuple[ + torch.Tensor, torch.Tensor +]: # logical_to_physical [num_logical_experts, num_ranks], replica_counts [num_logical_experts] + """构建单层逻辑 expert 到物理副本的映射""" + + logical_to_physical, replica_counts = _build_layer_maps( + redundant_expert_ids.unsqueeze(0), + num_logical_experts, + source_rank=source_rank, + node_world_size=node_world_size, + ) + return logical_to_physical.squeeze(0), replica_counts.squeeze(0) + + +def build_logical_to_physical_maps_for_layers( + redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] + num_logical_experts: int, # 逻辑 expert 的总数。 + source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 + node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 +) -> Tuple[ + torch.Tensor, # logical_to_physical, shape [num_layers, num_logical_experts, num_ranks] + torch.Tensor, # replica_counts, shape [num_layers, num_logical_experts] +]: + """为调用方传入的多个指定层构建逻辑 expert 到物理副本的 CPU int32 映射。 + + 第一维是层,不是请求或 token 的 batch,也不会自动处理模型中的其他层。 + 全局构建阶段第 0 列固定为主副本,其余列按 rank-major、slot-major 的稳定顺序 + 写入冗余副本。指定 source_rank 后,先筛选源节点内副本,再按 source_rank + 轮转候选前缀;最终返回映射的第 0 列只是第一个候选,不保证仍是主副本。 + """ + return _build_layer_maps( + redundant_expert_ids_by_layer, + num_logical_experts, + source_rank=source_rank, + node_world_size=node_world_size, + ) + + +def select_improving_placements( + expert_load: torch.Tensor, + current_placement: torch.Tensor, + candidate_placement: torch.Tensor, + *, + rebalance_gain_threshold: float, + expert_alignment: int | None = None, + node_world_size: int | None = None, +) -> Tuple[torch.Tensor, torch.Tensor, Dict[str, float | int], torch.Tensor, torch.Tensor]: + """Select better layers and return current/final rank loads without re-estimation.""" + if not 0.0 <= rebalance_gain_threshold <= 1.0: + raise ValueError("rebalance_gain_threshold must be between 0.0 and 1.0") + assert current_placement.shape == candidate_placement.shape + current_rank_load = _estimate_rank_load(expert_load, current_placement, expert_alignment, node_world_size) + candidate_rank_load = _estimate_rank_load(expert_load, candidate_placement, expert_alignment, node_world_size) + if expert_load.ndim == 2: + current_critical = current_rank_load.max(dim=1).values + candidate_critical = candidate_rank_load.max(dim=1).values + else: + current_critical = current_rank_load.max(dim=2).values.sum(dim=0) + candidate_critical = candidate_rank_load.max(dim=2).values.sum(dim=0) + # Each changed layer must reduce its own critical load. All selected + # changes must then collectively meet the configured model-level + # critical-load reduction threshold, avoiding low-gain migrations. + improved = candidate_critical < current_critical + selected = current_placement.clone() + selected[improved] = candidate_placement[improved] + if current_rank_load.ndim == 2: + selected_rank_load = torch.where(improved[:, None], candidate_rank_load, current_rank_load) + else: + selected_rank_load = torch.where(improved[None, :, None], candidate_rank_load, current_rank_load) + model_current_critical = current_critical.sum() + if expert_load.ndim == 2: + model_current_mean = current_rank_load.mean(dim=1).sum() + model_selected_critical = selected_rank_load.max(dim=1).values.sum() + model_selected_mean = selected_rank_load.mean(dim=1).sum() + else: + model_current_mean = current_rank_load.mean(dim=2).sum() + model_selected_critical = selected_rank_load.max(dim=2).values.sum() + model_selected_mean = selected_rank_load.mean(dim=2).sum() + model_ratio = model_current_critical / model_current_mean.clamp_min(1.0) + candidate_model_ratio = model_selected_critical / model_selected_mean.clamp_min(1.0) + candidate_rebalance_gain = (model_current_critical - model_selected_critical) / model_current_critical.clamp_min( + 1.0 + ) + metrics = { + "model_imbalance_ratio": float(model_ratio.item()), + "candidate_model_imbalance_ratio": float(candidate_model_ratio.item()), + "candidate_rebalance_gain": float(candidate_rebalance_gain.item()), + "candidate_changed_layer_count": int(improved.sum().item()), + } + if candidate_rebalance_gain >= rebalance_gain_threshold: + return selected, improved, metrics, current_rank_load, selected_rank_load + return ( + current_placement.clone(), + torch.zeros_like(improved), + metrics, + current_rank_load, + current_rank_load, + ) + + +def plan_redundant_experts( + expert_load: torch.Tensor, + num_ranks: int, + num_redundant_experts_per_rank: int, + expert_alignment: int | None = None, + node_world_size: int | None = None, + current_placement: torch.Tensor | None = None, + stickiness: float = 0.0, +) -> torch.Tensor: + """Plan replicas using source-node-local copies, with global fallback. + + With ``current_placement`` and a positive ``stickiness``, a candidate that + keeps an expert on its current rank receives a bonus of + ``stickiness * mean per-layer expert load``. This preserves rank + membership, not a particular redundant physical slot; target slots are + canonicalized against the current live rows before transfer and metadata + publication. A rank membership only changes when the move improves the + critical-load objective by more than that margin. + Without them the planning is bit-identical to the legacy behavior. + """ + assert expert_load.ndim in (2, 3, 4) + if expert_alignment is not None: + assert expert_alignment > 0 + use_legacy_topology_preference = expert_load.ndim < 4 + legacy_node_world_size = node_world_size if use_legacy_topology_preference else None + source_load, _squeeze_sample, node_world_size = _as_source_node_load(expert_load, num_ranks, node_world_size) + num_samples, num_layers, num_nodes, num_logical_experts = source_load.shape + assert num_logical_experts % num_ranks == 0 + assert num_redundant_experts_per_rank > 0 + num_experts_per_rank = num_logical_experts // num_ranks + num_redundant = num_ranks * num_redundant_experts_per_rank + assert num_redundant <= num_logical_experts * (num_ranks - 1) + + load = source_load.to(dtype=torch.float64, device="cpu") + placement = torch.full((num_layers, num_ranks, num_redundant_experts_per_rank), -1, dtype=torch.int64) + owner_rank = torch.arange(num_logical_experts, dtype=torch.int64) // num_experts_per_rank + if current_placement is not None: + assert tuple(current_placement.shape) == ( + num_layers, + num_ranks, + num_redundant_experts_per_rank, + ) + current_locations = _expert_locations(current_placement, num_logical_experts) + stickiness_scale = load.sum(dim=(0, 2, 3)) / num_logical_experts + else: + current_locations = None + stickiness_scale = None + + locations = _expert_locations(placement, num_logical_experts) + expert_rank = _expert_rank_load_all(load, locations, num_nodes, node_world_size, expert_alignment) + rank_load = expert_rank.sum(dim=2) + remaining_slots = torch.full((num_layers, num_ranks), num_redundant_experts_per_rank, dtype=torch.int64) + layer_indices = torch.arange(num_layers, dtype=torch.int64) + expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) + rank_nodes = ( + torch.arange(num_ranks, dtype=torch.int64) // legacy_node_world_size + if legacy_node_world_size is not None and legacy_node_world_size < num_ranks + else None + ) + + # Every iteration fills one slot per layer. Candidate expert evaluation + # is vectorized across all layers and logical experts, which keeps large + # GLM/Qwen planning comfortably on the CPU fast path. + for _ in range(num_redundant): + rank_order = torch.argsort(rank_load.sum(dim=0), dim=1, stable=True) + target_ranks = torch.full((num_layers,), -1, dtype=torch.int64) + legal = torch.zeros((num_layers, num_logical_experts), dtype=torch.bool) + for layer in range(num_layers): + for target_rank in rank_order[layer].tolist(): + if remaining_slots[layer, target_rank] == 0: + continue + candidate_legal = (owner_rank != target_rank) & ~locations[layer, :, target_rank] + # Legacy 2D/3D callers have no source-node axis. Retain the + # previous topology preference for that compatibility path; + # node-aware [S,L,N,E] planning uses only the exact load + # objective below. + if rank_nodes is not None: + existing_on_target_node = locations[layer, :, rank_nodes == rank_nodes[target_rank]].any(dim=1) + new_node_legal = candidate_legal & ~existing_on_target_node + if torch.any(new_node_legal): + candidate_legal = new_node_legal + if torch.any(candidate_legal): + target_ranks[layer] = target_rank + legal[layer] = candidate_legal + break + if torch.any(target_ranks < 0): + raise RuntimeError("EPLB planner found no valid redundant expert placement") + + candidate_locations = locations.clone() + candidate_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] = True + candidate_expert_rank = _expert_rank_load_all( + load, candidate_locations, num_nodes, node_world_size, expert_alignment + ) + candidate_rank_load = rank_load[:, :, None, :] - expert_rank + candidate_expert_rank + critical = candidate_rank_load.max(dim=3).values.sum(dim=0) + critical.masked_fill_(~legal, torch.inf) + if current_locations is not None: + # An expert already held by the target rank is retained unless + # another candidate beats it by more than the stickiness margin. + # This is rank membership, not physical-slot stickiness. Masked + # (inf) candidates stay masked: inf - x == inf. + keep = current_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] + critical = critical - stickiness * stickiness_scale[:, None] * keep + selected_experts = critical.argmin(dim=1) + if torch.isinf(critical[layer_indices, selected_experts]).any(): + raise RuntimeError("EPLB planner found no valid redundant expert placement") + + slots = num_redundant_experts_per_rank - remaining_slots[layer_indices, target_ranks] + placement[layer_indices, target_ranks, slots] = selected_experts + selected_next = candidate_expert_rank[:, layer_indices, selected_experts] + selected_old = expert_rank[:, layer_indices, selected_experts] + rank_load += selected_next - selected_old + expert_rank[:, layer_indices, selected_experts] = selected_next + locations[layer_indices, selected_experts, target_ranks] = True + remaining_slots[layer_indices, target_ranks] -= 1 + + assert torch.all(placement >= 0) + return placement + + +@dataclass(frozen=True, eq=False) +class _PhysicalExpertLayout: + """进程内按拓扑复用的只读物理 expert 布局;其中 Tensor 不得原地修改。""" + + num_logical_experts: int + num_ranks: int + num_physical_experts_per_rank: int + primary_physical_ids: torch.Tensor + redundant_physical_ids: torch.Tensor + + +@lru_cache(maxsize=8) +def _get_physical_expert_layout( + num_logical_experts: int, + num_ranks: int, + num_redundant_experts_per_rank: int, +) -> _PhysicalExpertLayout: + """返回按静态拓扑缓存的只读 CPU 物理 expert ID。""" + num_experts_per_rank = num_logical_experts // num_ranks + num_physical_experts_per_rank = num_experts_per_rank + num_redundant_experts_per_rank + expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) + primary_physical_ids = ( + (expert_ids // num_experts_per_rank) * num_physical_experts_per_rank + expert_ids % num_experts_per_rank + ).to(torch.int32) + ranks = torch.arange(num_ranks, dtype=torch.int64).repeat_interleave(num_redundant_experts_per_rank) + slots = torch.arange(num_redundant_experts_per_rank, dtype=torch.int64).repeat(num_ranks) + redundant_physical_ids = (ranks * num_physical_experts_per_rank + num_experts_per_rank + slots).to(torch.int32) + return _PhysicalExpertLayout( + num_logical_experts=num_logical_experts, + num_ranks=num_ranks, + num_physical_experts_per_rank=num_physical_experts_per_rank, + primary_physical_ids=primary_physical_ids, + redundant_physical_ids=redundant_physical_ids, + ) + + +def _build_global_replica_maps_for_layers( + redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] + layout: _PhysicalExpertLayout, +) -> Tuple[torch.Tensor, torch.Tensor]: + """为显式传入的多层冗余布局构建全局逻辑 expert id 到物理位置序列[rank, local_slot]的映射。 + + 输出: + logical_to_physical:CPU ``int32`` Tensor,形状为 + ``[num_layers, num_logical_experts, num_ranks]``。第 0 列固定为主 + 副本,后续列依次存放冗余副本,未使用位置为 ``-1``。 + replica_counts:CPU ``int32`` Tensor,形状为 + ``[num_layers, num_logical_experts]``。每个值包含主副本,并表示 + ``logical_to_physical`` 对应行中有效副本连续前缀的长度。 + + “全局映射”记录每层每个逻辑 expert 在所有 rank 上的主副本和冗余副本所对应的 + 物理 expert ID。第 0 列固定为主副本;冗余副本按所在 rank 从小到大、同一 rank + 内按槽位从小到大的顺序写入后续列。输入参数 + ``redundant_expert_ids_by_layer`` 已经指定每个 rank 的每个冗余槽位存放哪个逻辑 + expert,本函数只将该输入转换为顺序确定的映射,满足主副本优先、冗余槽位对应正确且有效副本连续排列的确定性结果。 + + """ + + num_layers = redundant_expert_ids_by_layer.shape[0] + num_logical_experts = layout.num_logical_experts + max_replicas = layout.num_ranks + redundant_ids = redundant_expert_ids_by_layer.to(dtype=torch.int64, device="cpu") + logical_to_physical = torch.full((num_layers, num_logical_experts, max_replicas), -1, dtype=torch.int32) + logical_to_physical[:, :, 0] = layout.primary_physical_ids + replica_counts = torch.ones((num_layers, num_logical_experts), dtype=torch.int32) + + flat_redundant_ids = redundant_ids.reshape(num_layers, -1) + if not flat_redundant_ids.numel(): + return logical_to_physical, replica_counts + + # 稳定排序保留 rank-major、slot-major 的历史顺序;第 0 列固定为主副本。 + sort_order = torch.argsort(flat_redundant_ids, dim=1, stable=True) + sorted_redundant_ids = flat_redundant_ids.gather(1, sort_order) + flat_positions = torch.arange(flat_redundant_ids.shape[1], dtype=torch.int64).unsqueeze(0) + group_starts = torch.where( + torch.cat( + ( + torch.ones((num_layers, 1), dtype=torch.bool), + sorted_redundant_ids[:, 1:] != sorted_redundant_ids[:, :-1], + ), + dim=1, + ), + flat_positions, + 0, + ) + replica_indices = flat_positions - torch.cummax(group_starts, dim=1).values + 1 + redundant_counts = torch.zeros((num_layers, num_logical_experts), dtype=torch.int32) + redundant_counts.scatter_add_( + 1, + flat_redundant_ids, + torch.ones_like(flat_redundant_ids, dtype=torch.int32), + ) + assert int(redundant_counts.max().item()) < max_replicas, "an expert can have at most one replica per rank" + replica_counts += redundant_counts + + layer_indices = torch.arange(num_layers, dtype=torch.int64).view(-1, 1).expand_as(sort_order) + redundant_physical_ids = layout.redundant_physical_ids.unsqueeze(0).expand_as(sort_order).gather(1, sort_order) + logical_to_physical[layer_indices, sorted_redundant_ids, replica_indices] = redundant_physical_ids + return logical_to_physical, replica_counts + + +def _select_source_node_replicas( + logical_to_physical: torch.Tensor, + replica_counts: torch.Tensor, + *, + source_rank: int, + node_world_size: int, + num_physical_experts_per_rank: int, + replica_positions: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + """按来源节点筛选全局候选副本,并将结果压缩为连续前缀。 + + ``logical_to_physical[layer, logical_expert]`` 的前 ``replica_counts`` 个位置 + 是有效副本,尾部 ``-1`` 表示没有副本;正常输入和输出都不会在有效前缀中 + 出现 ``-1``。连续的 ``node_world_size`` 个 rank 构成一个节点,来源节点为 + ``source_rank // node_world_size``。对每个逻辑 expert,若来源节点有副本, + 则只保留该节点的全部副本,远程副本(包括远程主副本)全部排除;否则保留 + 所有全局有效副本作为回退,避免候选为空。筛选可能选中原有效前缀中不连续的 + 位置,因此按原相对顺序将选中副本复制到新映射的连续前缀,其余位置填 ``-1``, + 并返回新的有效数量;不会原地修改输入。此函数只筛选和压缩候选集合,不做 + 负载优化、最终副本选择或按 source rank 轮转,轮转由后续函数完成。 + + 输出: + - ``compact_maps_by_layer``:与 ``logical_to_physical`` 同 shape + ``[num_layers, num_logical_experts, num_ranks]``,dtype/device 相同;每行 + 为筛选后副本的连续有效前缀,尾部为 ``-1``。 + - ``selected_counts_by_layer``:输入同 device 的 ``int32`` Tensor,shape 为 + ``[num_layers, num_logical_experts]``;每个值是对应输出行的有效前缀长度。 + """ + num_layers, num_logical_experts, _max_replicas = logical_to_physical.shape + source_node = source_rank // node_world_size + output_positions = replica_positions.view(1, 1, -1) + valid = output_positions < replica_counts.unsqueeze(-1) + local = valid & ( + torch.div( + logical_to_physical, + num_physical_experts_per_rank * node_world_size, + rounding_mode="floor", + ) + == source_node + ) + selected = torch.where(local.any(dim=2, keepdim=True), local, valid) + selected_counts_by_layer = selected.sum(dim=2, dtype=torch.int32) + + compact_maps_by_layer = torch.full_like(logical_to_physical, -1) + selected_positions = selected.cumsum(dim=2) - 1 + layers = torch.arange(num_layers, dtype=torch.int64).view(-1, 1, 1).expand_as(selected) + experts = torch.arange(num_logical_experts, dtype=torch.int64).view(1, -1, 1).expand_as(selected) + compact_maps_by_layer[layers[selected], experts[selected], selected_positions[selected]] = logical_to_physical[ + selected + ] + return compact_maps_by_layer, selected_counts_by_layer + + +def _rotate_selected_replicas( + compact_maps_by_layer: torch.Tensor, + selected_count_by_layer: torch.Tensor, + *, + source_rank: int, + replica_positions: torch.Tensor, +) -> torch.Tensor: + """根据 source_rank 循环调整每个逻辑 expert 的候选副本顺序,让不同源 rank 优先使用不同副本,同时保持候选副本集合和副本数量不变""" + output_positions = replica_positions.view(1, 1, -1) + selected_count64_by_layer = selected_count_by_layer.to(torch.int64).unsqueeze(-1) + rotation_by_layer = source_rank % selected_count64_by_layer + source_positions_by_layer = (output_positions + rotation_by_layer) % selected_count64_by_layer + maps_by_layer = compact_maps_by_layer.gather(2, source_positions_by_layer) + maps_by_layer.masked_fill_(output_positions >= selected_count64_by_layer, -1) + return maps_by_layer + + +def _build_layer_maps( + redundant_expert_ids_by_layer: torch.Tensor, + num_logical_experts: int, # 逻辑 expert 的总数。 + source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 + node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 +) -> Tuple[torch.Tensor, torch.Tensor]: + """构建调用方指定层的映射;全局构建、节点筛选和轮转均在此完成。""" + if redundant_expert_ids_by_layer.ndim != 3: + raise ValueError("redundant_expert_ids_by_layer must be [layers, ranks, num_redundant_experts_per_rank]") + num_ranks, num_redundant_experts_per_rank = redundant_expert_ids_by_layer.shape[1:] + assert num_logical_experts % num_ranks == 0 + layout = _get_physical_expert_layout(num_logical_experts, num_ranks, num_redundant_experts_per_rank) + logical_to_physical, replica_counts = _build_global_replica_maps_for_layers(redundant_expert_ids_by_layer, layout) + if source_rank is None: + return logical_to_physical, replica_counts + + assert node_world_size is not None + replica_positions = torch.arange(num_ranks, dtype=torch.int64) + compact_maps_by_layer, selected_counts_by_layer = _select_source_node_replicas( + logical_to_physical, + replica_counts, + source_rank=source_rank, + node_world_size=node_world_size, + num_physical_experts_per_rank=layout.num_physical_experts_per_rank, + replica_positions=replica_positions, + ) + return ( + _rotate_selected_replicas( + compact_maps_by_layer, + selected_counts_by_layer, + source_rank=source_rank, + replica_positions=replica_positions, + ), + selected_counts_by_layer, + ) + + +def _estimate_rank_load( + expert_load: torch.Tensor, + redundant_expert_ids: torch.Tensor, + expert_alignment: int | None = None, + node_world_size: int | None = None, +) -> torch.Tensor: + """Estimate runtime source-node-local routing load per physical expert. + + ``expert_load`` accepts the historic ``[layers, experts]`` and + ``[samples, layers, experts]`` forms, which are both one source node, and + the distributed ``[samples, layers, source_nodes, experts]`` form. Source + loads are kept separate until they are assigned to physical replicas, then + combined before applying the per-expert alignment used by DeepEP. + """ + source_load, squeeze_sample, node_world_size = _as_source_node_load( + expert_load, redundant_expert_ids.shape[1], node_world_size + ) + num_samples, num_layers, num_nodes, num_logical_experts = source_load.shape + assert redundant_expert_ids.ndim == 3 and redundant_expert_ids.shape[0] == num_layers + num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape[1:] + assert num_logical_experts % num_ranks == 0 + if expert_alignment is not None: + assert expert_alignment > 0 + + locations = _expert_locations(redundant_expert_ids, num_logical_experts) + route = _source_route(locations, num_nodes, node_world_size) + physical_load = torch.einsum("slne,lner->sler", source_load.to(torch.float64), route) + if expert_alignment is not None: + physical_load = torch.ceil(physical_load / expert_alignment) * expert_alignment + rank_load = physical_load.sum(dim=2) + return rank_load.squeeze(0) if squeeze_sample else rank_load + + +def _as_source_node_load( + expert_load: torch.Tensor, num_ranks: int, node_world_size: int | None +) -> Tuple[torch.Tensor, bool, int]: + """Normalize load to ``[samples, layers, source_nodes, experts]``.""" + assert expert_load.ndim in (2, 3, 4) + squeeze_sample = expert_load.ndim == 2 + if expert_load.ndim == 2: + source_load = expert_load.unsqueeze(0).unsqueeze(2) + elif expert_load.ndim == 3: + source_load = expert_load.unsqueeze(2) + else: + source_load = expert_load + num_nodes = source_load.shape[2] + # Historic 2D/3D loads represent one source node containing every rank. + if expert_load.ndim < 4: + return source_load, squeeze_sample, num_ranks + if node_world_size is None: + assert num_ranks % num_nodes == 0 + node_world_size = num_ranks // num_nodes + assert 0 < node_world_size <= num_ranks and num_ranks % node_world_size == 0 + assert num_nodes == num_ranks // node_world_size + return source_load, squeeze_sample, node_world_size + + +def _expert_locations(redundant_expert_ids: torch.Tensor, num_logical_experts: int) -> torch.Tensor: + """Return ``[layer, logical expert, rank]`` physical-copy occupancy.""" + num_layers, num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape + assert num_logical_experts % num_ranks == 0 + num_experts_per_rank = num_logical_experts // num_ranks + locations = torch.zeros( + (num_layers, num_logical_experts, num_ranks), + dtype=torch.bool, + device=redundant_expert_ids.device, + ) + expert_ids = torch.arange(num_logical_experts, device=locations.device) + owners = expert_ids // num_experts_per_rank + locations[:, expert_ids, owners] = True + layers = torch.arange(num_layers, device=locations.device)[:, None] + ranks = torch.arange(num_ranks, device=locations.device).repeat_interleave(num_redundant_experts_per_rank)[None, :] + redundant_ids = redundant_expert_ids.reshape(num_layers, -1) + valid = redundant_ids >= 0 + if torch.any(valid): + expanded_layers = layers.expand_as(redundant_ids) + expanded_ranks = ranks.expand_as(redundant_ids) + locations[ + expanded_layers[valid], + redundant_ids[valid], + expanded_ranks[valid], + ] = True + return locations + + +def _source_route(slots: torch.Tensor, num_nodes: int, node_world_size: int) -> torch.Tensor: + """Route each source node to its local copies, or all copies as fallback.""" + num_ranks = slots.shape[-1] + assert num_ranks % node_world_size == 0 and num_nodes == num_ranks // node_world_size + rank_nodes = torch.arange(num_ranks, device=slots.device) // node_world_size + source_nodes = torch.arange(num_nodes, device=slots.device) + copies = slots.unsqueeze(-3).expand(*slots.shape[:-2], num_nodes, *slots.shape[-2:]) + rank_node_shape = (1,) * slots.ndim + (num_ranks,) + source_node_shape = (1,) * (slots.ndim - 2) + (num_nodes, 1, 1) + local = copies & (rank_nodes.reshape(rank_node_shape) == source_nodes.reshape(source_node_shape)) + selected = torch.where(local.any(dim=-1, keepdim=True), local, copies) + return selected.to(torch.float64) / selected.sum(dim=-1, keepdim=True) + + +def _expert_rank_load_all( + source_load: torch.Tensor, + locations: torch.Tensor, + num_nodes: int, + node_world_size: int, + expert_alignment: int | None, +) -> torch.Tensor: + """Return aligned ``[samples, layers, expert, rank]`` contributions.""" + route = _source_route(locations, num_nodes, node_world_size) + physical_load = torch.einsum("slne,lner->sler", source_load.to(torch.float64), route) + if expert_alignment is not None: + physical_load = torch.ceil(physical_load / expert_alignment) * expert_alignment + return physical_load diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py new file mode 100644 index 0000000000..fc0f11e015 --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py @@ -0,0 +1,38 @@ +from dataclasses import dataclass +from typing import Optional + +import torch + + +@dataclass +class EPLBState: + num_redundant_experts_per_rank: int + initial_redundant_expert_ids_by_rank: torch.Tensor + logical_to_physical_map: torch.Tensor + logical_replica_count: torch.Tensor + route_counter: torch.Tensor + recording: bool = False + recorded_sample_count: int = 0 + + def next_sample_index(self) -> int: + if not self.recording: + return 0 + sample_index = self.recorded_sample_count % self.route_counter.shape[0] + self.recorded_sample_count += 1 + return sample_index + + +@dataclass(frozen=True) +class ExpertParallelState: + num_logical_experts: int + world_size: int + eplb: Optional[EPLBState] = None + + @property + def num_primary_experts_per_rank(self) -> int: + return self.num_logical_experts // self.world_size + + @property + def num_total_physical_experts(self) -> int: + num_redundant_experts_per_rank = 0 if self.eplb is None else self.eplb.num_redundant_experts_per_rank + return self.num_logical_experts + self.world_size * num_redundant_experts_per_rank 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 6fc81dd6f9..225be834b2 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 @@ -8,10 +8,19 @@ get_col_slice_mixin, SliceMixinTpl, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import select_fuse_moe_impl +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import create_fuse_moe_impl +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + EPLBState, + ExpertParallelState, +) from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_initial_redundant_expert_ids, + build_logical_to_physical_map, +) from lightllm.common.quantization.quantize_method import QuantizationMethod -from lightllm.utils.dist_utils import get_global_world_size, get_global_rank +from lightllm.utils.envs_utils import get_env_start_args, get_prefill_eplb_step_interval +from lightllm.utils.dist_utils import get_global_world_size, get_global_rank, get_node_world_size from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) @@ -55,15 +64,14 @@ def __init__( self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self._init_config(network_config) - self._init_redundancy_expert_params() - self._init_parallel_params() - self.fuse_moe_impl = select_fuse_moe_impl(self.quant_method, self.enable_ep_moe)( + self._init_expert_parallel_state() + self._init_weight_partition() + self.fuse_moe_impl = create_fuse_moe_impl( n_routed_experts=self.n_routed_experts, num_fused_shared_experts=self.num_fused_shared_experts, routed_scaling_factor=self.routed_scaling_factor, quant_method=self.quant_method, - redundancy_expert_num=self.redundancy_expert_num, - routed_expert_counter_tensor=self.routed_expert_counter_tensor, + expert_parallel_state=self.expert_parallel_state, ) self.lock = threading.Lock() self._create_weight() @@ -77,12 +85,50 @@ def _init_config(self, network_config: Dict[str, Any]): self.routed_scaling_factor = network_config.get("routed_scaling_factor", 1.0) self.scoring_func = network_config.get("scoring_func", "softmax") - def _init_redundancy_expert_params(self): - self.redundancy_expert_num = 0 - self.redundancy_expert_ids = [] - self.routed_expert_counter_tensor = torch.zeros((self.n_routed_experts,), dtype=torch.int64, device="cuda") + def _init_expert_parallel_state(self): + args = get_env_start_args() + self.expert_parallel_state: Optional[ExpertParallelState] = None + # Initial placement metadata is used only while loading checkpoint rows. + self._initial_redundant_expert_ids = [] + self._initial_redundant_expert_idx_to_local_idx = {} + eplb = None + if args.enable_prefill_eplb: + num_redundant_experts_per_rank = args.eplb_num_redundant_experts_per_rank + all_initial_ids = build_initial_redundant_expert_ids( + self.n_routed_experts, + self.global_world_size, + num_redundant_experts_per_rank, + ) + self._initial_redundant_expert_ids = all_initial_ids[self.global_rank_].tolist() + logical_to_physical, logical_replica_count = build_logical_to_physical_map( + all_initial_ids, + self.n_routed_experts, + source_rank=self.global_rank_, + node_world_size=get_node_world_size(), + ) + # route_counter 每次 prefill dispatch 记录一行。初始阶段连续采样 + # step_interval 个 manager step,兼顾micro batch overlap的两次 dispatch,因此容量设为 + # 2 * step_interval。稳定阶段复用该环形缓冲区,但只把当前短采样窗口内 + # 实际记录的最近行传给 planner,不复制整个缓冲区。 + eplb = EPLBState( + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + initial_redundant_expert_ids_by_rank=all_initial_ids, + logical_to_physical_map=logical_to_physical.cuda(), + logical_replica_count=logical_replica_count.cuda(), + route_counter=torch.zeros( + (2 * get_prefill_eplb_step_interval(), self.n_routed_experts), + dtype=torch.int64, + device="cuda", + ), + ) + if self.enable_ep_moe: + self.expert_parallel_state = ExpertParallelState( + num_logical_experts=self.n_routed_experts, + world_size=self.global_world_size, + eplb=eplb, + ) - def _init_parallel_params(self): + def _init_weight_partition(self): if self.enable_ep_moe: self.tp_rank_ = 0 self.tp_world_size_ = 1 @@ -96,27 +142,26 @@ def _init_parallel_params(self): self.split_inter_size = self.moe_intermediate_size // self.tp_world_size_ if self.enable_ep_moe: assert self.num_fused_shared_experts == 0, "num_fused_shared_experts must be 0 when enable_ep_moe" + eplb = self.expert_parallel_state.eplb + num_redundant_experts_per_rank = 0 if eplb is None else eplb.num_redundant_experts_per_rank logger.debug( f"global_rank {self.global_rank_} layerindex {self.layer_num_} " - f"redundancy_expertids: {self.redundancy_expert_ids}" - ) - self.local_n_routed_experts = self.n_routed_experts // self.global_world_size + self.redundancy_expert_num - n_experts_per_rank = self.n_routed_experts // self.global_world_size - start_expert_id = self.global_rank_ * n_experts_per_rank - self.local_expert_ids = ( - list(range(start_expert_id, start_expert_id + n_experts_per_rank)) + self.redundancy_expert_ids + f"initial_redundant_expert_ids: {self._initial_redundant_expert_ids}" ) + num_primary_experts_per_rank = self.expert_parallel_state.num_primary_experts_per_rank + self.local_n_routed_experts = num_primary_experts_per_rank + num_redundant_experts_per_rank + start_expert_id = self.global_rank_ * num_primary_experts_per_rank self.expert_idx_to_local_idx = { - expert_idx: expert_idx - start_expert_id for expert_idx in self.local_expert_ids[:n_experts_per_rank] + expert_idx: expert_idx - start_expert_id + for expert_idx in range(start_expert_id, start_expert_id + num_primary_experts_per_rank) } - self.redundancy_expert_idx_to_local_idx = { - redundancy_expert_idx: n_experts_per_rank + i - for (i, redundancy_expert_idx) in enumerate(self.redundancy_expert_ids) + self._initial_redundant_expert_idx_to_local_idx = { + redundant_expert_idx: num_primary_experts_per_rank + i + for (i, redundant_expert_idx) in enumerate(self._initial_redundant_expert_ids) } else: self.local_expert_ids = list(range(self.n_routed_experts + self.num_fused_shared_experts)) self.expert_idx_to_local_idx = {expert_idx: i for (i, expert_idx) in enumerate(self.local_expert_ids)} - self.rexpert_idx_to_local_idx = {} def experts( self, @@ -274,8 +319,8 @@ def load_hf_weights(self, weights): self._load_e_score_correction_bias(weights) self._load_per_expert_scale(weights) self._load_weight(self.expert_idx_to_local_idx, weights) - if self.redundancy_expert_num > 0: - self._load_weight(self.redundancy_expert_idx_to_local_idx, weights) + if self._initial_redundant_expert_idx_to_local_idx: + self._load_weight(self._initial_redundant_expert_idx_to_local_idx, weights) def verify_load(self): weight_load_ok = all(all(_weight_pack.load_ok) for _weight_pack in self.w1_list + self.w2_list + self.w3_list) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 67bb90e4ef..c00a35f600 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -2,13 +2,36 @@ from .triton_impl import FuseMoeTriton from .marlin_impl import FuseMoeMarlin from .deepgemm_impl import FuseMoeDeepGEMM +from ..expert_parallel_state import ExpertParallelState -def select_fuse_moe_impl(quant_method: QuantizationMethod, enable_ep_moe: bool): - if enable_ep_moe: - return FuseMoeDeepGEMM +def create_fuse_moe_impl( + *, + n_routed_experts: int, + num_fused_shared_experts: int, + routed_scaling_factor: float, + quant_method: QuantizationMethod, + expert_parallel_state: ExpertParallelState | None = None, +): + if expert_parallel_state is not None: + return FuseMoeDeepGEMM( + n_routed_experts=n_routed_experts, + num_fused_shared_experts=num_fused_shared_experts, + routed_scaling_factor=routed_scaling_factor, + quant_method=quant_method, + expert_parallel_state=expert_parallel_state, + ) if quant_method.method_name == "awq_marlin": - return FuseMoeMarlin - else: - return FuseMoeTriton + return FuseMoeMarlin( + n_routed_experts=n_routed_experts, + num_fused_shared_experts=num_fused_shared_experts, + routed_scaling_factor=routed_scaling_factor, + quant_method=quant_method, + ) + return FuseMoeTriton( + n_routed_experts=n_routed_experts, + num_fused_shared_experts=num_fused_shared_experts, + routed_scaling_factor=routed_scaling_factor, + quant_method=quant_method, + ) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py index 9b6c42af79..35e872df10 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py @@ -1,49 +1,25 @@ import torch -from abc import abstractmethod -from typing import Callable, Optional +from abc import ABC, abstractmethod +from typing import Callable, Optional, Tuple from lightllm.common.quantization.quantize_method import ( WeightPack, QuantizationMethod, ) -from lightllm.utils.dist_utils import ( - get_global_rank, - get_global_world_size, -) -class FuseMoeBaseImpl: +class FuseMoeBaseImpl(ABC): def __init__( self, n_routed_experts: int, num_fused_shared_experts: int, routed_scaling_factor: float, quant_method: QuantizationMethod, - redundancy_expert_num: int, - routed_expert_counter_tensor: torch.Tensor, ): self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self.routed_scaling_factor = routed_scaling_factor self.quant_method = quant_method - self.global_rank_ = get_global_rank() - self.global_world_size_ = get_global_world_size() - self.ep_n_routed_experts = self.n_routed_experts // self.global_world_size_ - self.total_expert_num_contain_redundancy = ( - self.n_routed_experts + redundancy_expert_num * self.global_world_size_ - ) - - # redundancy expert related - self.redundancy_expert_num = redundancy_expert_num - self.routed_expert_counter_tensor = routed_expert_counter_tensor - # workspace for kernel optimization - self.workspace = self.create_workspace() - - @abstractmethod - def create_workspace(self): - pass - - @abstractmethod def __call__( self, input_tensor: torch.Tensor, @@ -63,5 +39,62 @@ def __call__( per_expert_scale: Optional[torch.Tensor] = None, # Qwen3.5 uses this gate to control fused shared expert aggregation weights. shared_expert_gate: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + topk_weights, topk_ids, origin_topk_ids = self._select_experts( + input_tensor=input_tensor, + router_logits=router_logits, + correction_bias=correction_bias, + top_k=top_k, + renormalize=renormalize, + use_grouped_topk=use_grouped_topk, + topk_group=topk_group, + num_expert_group=num_expert_group, + scoring_func=scoring_func, + per_expert_scale=per_expert_scale, + shared_expert_gate=shared_expert_gate, + is_prefill=is_prefill, + preserve_logical_ids=moe_capture_callback is not None, + ) + if moe_capture_callback is not None: + moe_capture_callback(origin_topk_ids) + return self._fused_experts( + input_tensor=input_tensor, + w13=w13, + w2=w2, + topk_weights=topk_weights, + topk_ids=topk_ids, + router_logits=router_logits, + is_prefill=is_prefill, + ) + + @abstractmethod + def _select_experts( + self, + input_tensor: torch.Tensor, + router_logits: torch.Tensor, + correction_bias: Optional[torch.Tensor], + top_k: int, + renormalize: bool, + use_grouped_topk: bool, + topk_group: int, + num_expert_group: int, + scoring_func: str, + per_expert_scale: Optional[torch.Tensor] = None, + shared_expert_gate: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + preserve_logical_ids: bool = False, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + pass + + @abstractmethod + def _fused_experts( + self, + input_tensor: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + router_logits: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, ) -> torch.Tensor: pass 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 6fabd3304c..f234f09896 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 @@ -1,6 +1,7 @@ import torch from typing import Optional, Tuple, Any -from .triton_impl import FuseMoeTriton +from .base_impl import FuseMoeBaseImpl +from ..expert_parallel_state import ExpertParallelState from lightllm.distributed import dist_group_manager from lightllm.common.quantization.quantize_method import WeightPack from lightllm.utils.envs_utils import ( @@ -18,10 +19,13 @@ from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType -class FuseMoeDeepGEMM(FuseMoeTriton): - def __init__(self, *args, **kwargs): +class FuseMoeDeepGEMM(FuseMoeBaseImpl): + def __init__(self, *args, expert_parallel_state: ExpertParallelState, **kwargs): super().__init__(*args, **kwargs) + self.expert_parallel_state = expert_parallel_state + self.eplb = expert_parallel_state.eplb self.ep_balance_counters = None + self._primary_weight_pack_cache = {} def _select_experts( self, @@ -36,27 +40,55 @@ def _select_experts( scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, shared_expert_gate: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + preserve_logical_ids: bool = False, ): - """Select experts and return topk weights and ids.""" + """选择 expert;EPLB prefill 统一由融合路径返回 physical ID。""" assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" - from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts + eplb = self.eplb + eplb_active = eplb is not None + if is_prefill is True and eplb_active: + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb - topk_weights, topk_ids = select_experts( - hidden_states=input_tensor, - router_logits=router_logits, - correction_bias=correction_bias, - use_grouped_topk=use_grouped_topk, - top_k=top_k, - renormalize=renormalize, - topk_group=topk_group, - num_expert_group=num_expert_group, - scoring_func=scoring_func, - ) + group_score_topk_num = 2 if topk_group == 4 and num_expert_group == 8 and top_k == 8 else 1 + topk_weights, topk_ids, logical_topk_ids = triton_grouped_topk_eplb( + hidden_states=input_tensor, + gating_output=router_logits, + correction_bias=correction_bias, + topk=top_k, + renormalize=renormalize, + num_expert_group=num_expert_group, + topk_group=topk_group, + scoring_func=scoring_func, + logical_to_physical_map=eplb.logical_to_physical_map, + logical_replica_count=eplb.logical_replica_count, + expert_counter=eplb.route_counter, + sample_index=eplb.next_sample_index(), + record_load=eplb.recording, + use_grouped_topk=use_grouped_topk, + return_logical_ids=preserve_logical_ids, + group_score_used_topk_num=group_score_topk_num, + ) + origin_topk_ids = logical_topk_ids if logical_topk_ids is not None else topk_ids + else: + from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts + + topk_weights, topk_ids = select_experts( + hidden_states=input_tensor, + router_logits=router_logits, + correction_bias=correction_bias, + use_grouped_topk=use_grouped_topk, + top_k=top_k, + renormalize=renormalize, + topk_group=topk_group, + num_expert_group=num_expert_group, + scoring_func=scoring_func, + ) + if per_expert_scale is not None: + topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype) + origin_topk_ids = topk_ids if self.routed_scaling_factor != 1.0: topk_weights.mul_(self.routed_scaling_factor) - if per_expert_scale is not None: - topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype) - origin_topk_ids = topk_ids return topk_weights, topk_ids, origin_topk_ids def _fused_experts( @@ -69,13 +101,20 @@ def _fused_experts( router_logits: Optional[torch.Tensor] = None, is_prefill: Optional[bool] = None, ): + if is_prefill is False: + w13 = self._primary_weight_pack(w13) + w2 = self._primary_weight_pack(w2) + num_experts = self.n_routed_experts + else: + num_experts = self.expert_parallel_state.num_total_physical_experts + output = fused_experts( hidden_states=input_tensor, w13=w13, w2=w2, topk_weights=topk_weights, topk_idx=topk_ids.to(torch.long), - num_experts=self.total_expert_num_contain_redundancy, # number of all experts contain redundancy + num_experts=num_experts, quant_method=self.quant_method, is_prefill=is_prefill, previous_event=None, # for overlap @@ -105,6 +144,7 @@ def low_latency_dispatch( topk_group=topk_group, num_expert_group=n_group, scoring_func=scoring_func, + is_prefill=False, ) topk_idx = topk_idx.to(torch.long) @@ -114,7 +154,8 @@ def low_latency_dispatch( topk_idx=topk_idx, x=hidden_states, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, - num_experts=self.total_expert_num_contain_redundancy, + # decode 与 EPLB 的物理冗余行刻意隔离:DeepEP 使用原始 logical expert ID。 + num_experts=self.n_routed_experts, use_fp8=use_fp8_w8a8, async_finish=False, return_recv_hook=True, @@ -144,6 +185,7 @@ def select_experts_and_quant_input( topk_group=topk_group, num_expert_group=n_group, scoring_func=scoring_func, + is_prefill=True, ) qinput_tensor = quantize_fused_experts_input(hidden_states, w13, self.quant_method) return topk_weights, topk_idx.to(torch.long), qinput_tensor @@ -161,7 +203,7 @@ def dispatch( qinput_tensor, topk_idx=topk_idx, topk_weights=topk_weights, - num_experts=self.total_expert_num_contain_redundancy, + num_experts=self.expert_parallel_state.num_total_physical_experts, num_max_tokens_per_rank=num_max_tokens_per_rank, expert_alignment=128, num_sms=get_ep_num_sms(), @@ -203,6 +245,7 @@ def masked_group_gemm( dtype: torch.dtype, expected_m: int, ): + w13, w2 = self._primary_weight_pack(w13), self._primary_weight_pack(w2) w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale return masked_group_gemm( @@ -300,3 +343,30 @@ def hook(): event.current_stream_wait() return combined_x, hook + + def _primary_weight_pack(self, weight_pack: WeightPack) -> WeightPack: + """返回所有 decode 路径使用的缓存本地主副本视图。""" + if self.eplb is None: + return weight_pack + cache = getattr(self, "_primary_weight_pack_cache", None) + if cache is None: + cache = self._primary_weight_pack_cache = {} + cache_key = id(weight_pack) + primary = cache.get(cache_key) + if primary is None: + num_primary_experts_per_rank = self.expert_parallel_state.num_primary_experts_per_rank + primary = WeightPack( + weight=weight_pack.weight[:num_primary_experts_per_rank], + weight_scale=( + weight_pack.weight_scale[:num_primary_experts_per_rank] + if weight_pack.weight_scale is not None + else None + ), + weight_zero_point=( + getattr(weight_pack, "weight_zero_point", None)[:num_primary_experts_per_rank] + if getattr(weight_pack, "weight_zero_point", None) is not None + else None + ), + ) + cache[cache_key] = primary + return primary diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py index 0094b09b1c..2ee57fe916 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py @@ -11,6 +11,10 @@ class FuseMoeMarlin(FuseMoeTriton): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.workspace = self.create_workspace() + def create_workspace(self): from lightllm.utils.vllm_utils import HAS_VLLM diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index abe0112004..c5f62ae946 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -1,32 +1,10 @@ import torch -from typing import Callable, Optional +from typing import Optional from lightllm.common.quantization.no_quant import WeightPack -from lightllm.common.quantization.quantize_method import QuantizationMethod from .base_impl import FuseMoeBaseImpl class FuseMoeTriton(FuseMoeBaseImpl): - def __init__( - self, - n_routed_experts: int, - num_fused_shared_experts: int, - routed_scaling_factor: float, - quant_method: QuantizationMethod, - redundancy_expert_num: int, - routed_expert_counter_tensor: torch.Tensor, - ): - super().__init__( - n_routed_experts=n_routed_experts, - num_fused_shared_experts=num_fused_shared_experts, - routed_scaling_factor=routed_scaling_factor, - quant_method=quant_method, - redundancy_expert_num=redundancy_expert_num, - routed_expert_counter_tensor=routed_expert_counter_tensor, - ) - - def create_workspace(self): - return None - def _select_experts( self, input_tensor: torch.Tensor, @@ -40,6 +18,8 @@ def _select_experts( scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, shared_expert_gate: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + preserve_logical_ids: bool = False, ): """Select experts and return topk weights and ids.""" from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts @@ -102,50 +82,3 @@ def _fused_experts( w2_scale=w2_scale, ) return input_tensor - - def __call__( - self, - input_tensor: torch.Tensor, - router_logits: torch.Tensor, - w13: WeightPack, - w2: WeightPack, - correction_bias: Optional[torch.Tensor], - scoring_func: str, - top_k: int, - renormalize: bool, - use_grouped_topk: bool, - topk_group: int, - num_expert_group: int, - is_prefill: Optional[bool] = None, - # Callback to capture MoE topk expert ids (routed experts metadata). - moe_capture_callback: Optional[Callable[[torch.Tensor], None]] = None, - per_expert_scale: Optional[torch.Tensor] = None, - shared_expert_gate: Optional[torch.Tensor] = None, - ): - topk_weights, topk_ids, origin_topk_ids = self._select_experts( - input_tensor=input_tensor, - router_logits=router_logits, - correction_bias=correction_bias, - top_k=top_k, - renormalize=renormalize, - use_grouped_topk=use_grouped_topk, - topk_group=topk_group, - num_expert_group=num_expert_group, - scoring_func=scoring_func, - per_expert_scale=per_expert_scale, - shared_expert_gate=shared_expert_gate, - ) - - if moe_capture_callback is not None: - moe_capture_callback(origin_topk_ids) - - output = self._fused_experts( - input_tensor=input_tensor, - w13=w13, - w2=w2, - topk_weights=topk_weights, - topk_ids=topk_ids, - router_logits=router_logits, - is_prefill=is_prefill, - ) - return output diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_kernels.py b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_kernels.py new file mode 100644 index 0000000000..15462dff0b --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_kernels.py @@ -0,0 +1,55 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def eplb_replica_index(token_index, logical_id, replica_count): + """Choose a replica with independent phases for a token's top-k experts.""" + token_hash = token_index.to(tl.uint32) * 2654435769 + expert_hash = logical_id.to(tl.uint32) * 2246822519 + return (token_hash + expert_hash) % replica_count.to(tl.uint32) + + +@triton.jit +def _eplb_push_copy_kernel( + src_ptrs_ptr, + dst_ptrs_ptr, + bytes_per_descriptor, + BLOCK_SIZE: tl.constexpr, + ITEMS_PER_PROGRAM: tl.constexpr, +): + descriptor_index = tl.program_id(1) + offsets = tl.program_id(0) * (BLOCK_SIZE * ITEMS_PER_PROGRAM) + tl.arange(0, BLOCK_SIZE) + src_ptr = tl.load(src_ptrs_ptr + descriptor_index).to(tl.pointer_type(tl.uint64)) + dst_ptr = tl.load(dst_ptrs_ptr + descriptor_index).to(tl.pointer_type(tl.uint64)) + word_count = bytes_per_descriptor // 8 + for item_index in tl.static_range(0, ITEMS_PER_PROGRAM): + word_offsets = offsets + item_index * BLOCK_SIZE + mask = word_offsets < word_count + values = tl.load(src_ptr + word_offsets, mask=mask, cache_modifier=".cg") + tl.store(dst_ptr + word_offsets, values, mask=mask, cache_modifier=".cs") + + +@torch.no_grad() +def eplb_push_copy(src_ptrs: torch.Tensor, dst_ptrs: torch.Tensor, bytes_per_descriptor: int) -> None: + """Copy 16-byte-aligned expert rows from source to destination pointers.""" + if bytes_per_descriptor <= 64 * 1024: + block_size = 128 + num_warps = 4 + elif bytes_per_descriptor >= 4 * 1024 * 1024 and src_ptrs.numel() > 1: + block_size = 512 + num_warps = 8 + else: + block_size = 256 + num_warps = 4 + items_per_program = 4 + words_per_program = block_size * items_per_program + _eplb_push_copy_kernel[(triton.cdiv(bytes_per_descriptor // 8, words_per_program), src_ptrs.numel())]( + src_ptrs, + dst_ptrs, + bytes_per_descriptor, + BLOCK_SIZE=block_size, + ITEMS_PER_PROGRAM=items_per_program, + num_warps=num_warps, + ) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py index fb0323cd4b..eabe6a6311 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py @@ -4,6 +4,8 @@ import triton.language as tl from triton.language.standard import _log2, sum, zeros_like +from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_kernels import eplb_replica_index + @triton.jit def _compare_and_swap(x, x_1, ids, flip, i: tl.core.constexpr, n_dims: tl.core.constexpr): @@ -202,6 +204,172 @@ def grouped_topk_kernel( return +@triton.jit +def grouped_topk_eplb_kernel( + gating_output_ptr, + gating_output_stride_m, + gating_output_stride_n, + correction_bias_ptr, + out_topk_weights, + out_topk_weights_stride_m, + out_topk_weights_stride_n, + out_topk_ids, + out_topk_ids_stride_m, + out_topk_ids_stride_n, + out_logical_ids, + out_logical_ids_stride_m, + out_logical_ids_stride_n, + logical_to_physical_ptr, + logical_replica_count_ptr, + expert_counter_ptr, + sample_index, + group_num, + group_expert_num, + total_expert_num, + group_topk_num, + IS_SIGMOID: tl.constexpr, + USE_GROUPED_TOPK: tl.constexpr, + HAS_CORRECTION_BIAS: tl.constexpr, + RETURN_LOGICAL_IDS: tl.constexpr, + EXPERT_GROUP_NUM: tl.constexpr, + EXPERT_GROUP_SIZE: tl.constexpr, + TOPK_NUM: tl.constexpr, + TOPK_BLOCK_SIZE: tl.constexpr, + RENORMALIZE: tl.constexpr, + GROUP_SCORE_USED_TOPK_NUM: tl.constexpr, + COUNTER_NUM_EXPERTS: tl.constexpr, + MAP_SLOTS: tl.constexpr, + RECORD_LOAD: tl.constexpr, + SINGLE_TOKEN: tl.constexpr, +): + """Grouped top-k, EPLB accounting, and replica mapping without a global score workspace.""" + token_index = tl.program_id(axis=0) + offs_group = tl.arange(0, EXPERT_GROUP_NUM) + offs_group_v = tl.arange(0, EXPERT_GROUP_SIZE) + logical_ids = offs_group[:, None] * group_expert_num + offs_group_v[None, :] + valid_expert = ( + (offs_group < group_num)[:, None] + & (offs_group_v < group_expert_num)[None, :] + & (logical_ids < total_expert_num) + ) + hidden_states = tl.load( + gating_output_ptr + token_index * gating_output_stride_m + logical_ids * gating_output_stride_n, + mask=valid_expert, + other=-float("inf"), + ).to(tl.float32) + + if IS_SIGMOID: + old_scores = tl.sigmoid(hidden_states) + else: + group_max = tl.max(hidden_states, axis=1) + global_max = tl.max(group_max, axis=0) + numerators = tl.where(valid_expert, tl.exp(hidden_states - global_max), 0.0) + denominator = tl.sum(tl.sum(numerators, axis=1), axis=0) + old_scores = numerators / denominator + + if HAS_CORRECTION_BIAS: + correction_bias = tl.load(correction_bias_ptr + logical_ids, mask=valid_expert, other=0.0) + scores = tl.where(valid_expert, old_scores + correction_bias, -float("inf")) + else: + scores = tl.where(valid_expert, old_scores, -float("inf")) + + if USE_GROUPED_TOPK: + if GROUP_SCORE_USED_TOPK_NUM == 1: + group_value = tl.max(scores, axis=1) + elif GROUP_SCORE_USED_TOPK_NUM == 2: + first_score, first_index = tl.max(scores, axis=1, return_indices=True) + second_score = tl.max( + tl.where(offs_group_v[None, :] == first_index[:, None], -float("inf"), scores), + axis=1, + ) + group_value = first_score + second_score + else: + sorted_group_scores = tl.sort(scores, dim=1, descending=True) + group_value = tl.sum( + tl.where(offs_group_v[None, :] < GROUP_SCORE_USED_TOPK_NUM, sorted_group_scores, 0.0), + axis=1, + ) + + if EXPERT_GROUP_NUM > 1: + sorted_group_value = tl.sort(group_value, descending=True) + else: + sorted_group_value = group_value + group_topk_value = tl.sum(tl.where(offs_group == group_topk_num - 1, sorted_group_value, 0.0)) + candidate_scores = tl.where( + (group_value >= group_topk_value)[:, None] & valid_expert, + scores, + -float("inf"), + ) + else: + candidate_scores = tl.where(valid_expert, old_scores, -float("inf")) + + sort_block_size: tl.constexpr = EXPERT_GROUP_NUM * EXPERT_GROUP_SIZE + flat_offsets = tl.arange(0, sort_block_size) + candidate_scores = tl.reshape(candidate_scores, (sort_block_size,)) + topk_offsets = tl.arange(0, TOPK_BLOCK_SIZE) + selected_weights = tl.zeros((TOPK_BLOCK_SIZE,), tl.float32) + selected_logical_ids = tl.zeros((TOPK_BLOCK_SIZE,), tl.int32) + sum_scores = 0.0 + for topk_index in range(TOPK_NUM): + selected_offset = tl.argmax(candidate_scores, axis=0) + selected_group = selected_offset // EXPERT_GROUP_SIZE + selected_group_offset = selected_offset % EXPERT_GROUP_SIZE + selected_logical_id = selected_group * group_expert_num + selected_group_offset + selected_hidden_state = tl.load( + gating_output_ptr + token_index * gating_output_stride_m + selected_logical_id * gating_output_stride_n + ).to(tl.float32) + if IS_SIGMOID: + selected_weight = tl.sigmoid(selected_hidden_state) + else: + selected_weight = tl.exp(selected_hidden_state - global_max) / denominator + sum_scores += selected_weight + topk_lane = topk_offsets == topk_index + selected_weights = tl.where(topk_lane, selected_weight, selected_weights) + selected_logical_ids = tl.where(topk_lane, selected_logical_id, selected_logical_ids) + candidate_scores = tl.where(flat_offsets == selected_offset, -float("inf"), candidate_scores) + + topk_mask = topk_offsets < TOPK_NUM + if RECORD_LOAD: + tl.atomic_add( + expert_counter_ptr + sample_index * COUNTER_NUM_EXPERTS + selected_logical_ids, + 1, + mask=topk_mask, + sem="relaxed", + ) + if SINGLE_TOKEN: + replica_indices = tl.zeros((TOPK_BLOCK_SIZE,), tl.int32) + else: + replica_counts = tl.load( + logical_replica_count_ptr + selected_logical_ids, + mask=topk_mask, + other=1, + ) + replica_indices = eplb_replica_index(token_index, selected_logical_ids, replica_counts) + selected_physical_ids = tl.load( + logical_to_physical_ptr + selected_logical_ids * MAP_SLOTS + replica_indices, + mask=topk_mask, + other=-1, + ) + if RENORMALIZE: + selected_weights /= sum_scores + tl.store( + out_topk_weights + token_index * out_topk_weights_stride_m + topk_offsets * out_topk_weights_stride_n, + selected_weights, + mask=topk_mask, + ) + tl.store( + out_topk_ids + token_index * out_topk_ids_stride_m + topk_offsets * out_topk_ids_stride_n, + selected_physical_ids, + mask=topk_mask, + ) + if RETURN_LOGICAL_IDS: + tl.store( + out_logical_ids + token_index * out_logical_ids_stride_m + topk_offsets * out_logical_ids_stride_n, + selected_logical_ids, + mask=topk_mask, + ) + + def triton_grouped_topk( hidden_states: torch.Tensor, gating_output: torch.Tensor, @@ -263,3 +431,81 @@ def triton_grouped_topk( num_stages=1, ) return out_topk_weights, out_topk_ids + + +def triton_grouped_topk_eplb( + hidden_states: torch.Tensor, + gating_output: torch.Tensor, + correction_bias: torch.Tensor, + topk: int, + renormalize: bool, + num_expert_group: int, + topk_group: int, + scoring_func: str, + logical_to_physical_map: torch.Tensor, + logical_replica_count: torch.Tensor, + expert_counter: torch.Tensor, + sample_index: int, + record_load: bool, + use_grouped_topk: bool, + return_logical_ids: bool = False, + group_score_used_topk_num: int = 2, +): + """Fused EPLB prefill top-k returning physical IDs and optional logical IDs.""" + token_num, total_expert_num = gating_output.shape + out_topk_weights = torch.empty((token_num, topk), dtype=torch.float32, device=gating_output.device) + out_topk_ids = torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) + out_logical_ids = ( + torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) if return_logical_ids else None + ) + if token_num == 0: + return out_topk_weights, out_topk_ids, out_logical_ids + if use_grouped_topk: + assert total_expert_num % num_expert_group == 0 + group_num = num_expert_group + group_expert_num = total_expert_num // num_expert_group + group_topk_num = topk_group + else: + group_num = 1 + group_expert_num = total_expert_num + group_topk_num = 1 + expert_group_num = triton.next_power_of_2(group_num) + expert_group_size = triton.next_power_of_2(group_expert_num) + sort_block_size = expert_group_num * expert_group_size + num_warps = min(max(1, sort_block_size // 256), 8) + grouped_topk_eplb_kernel[(token_num,)]( + gating_output, + *gating_output.stride(), + correction_bias, + out_topk_weights, + *out_topk_weights.stride(), + out_topk_ids, + *out_topk_ids.stride(), + out_logical_ids if out_logical_ids is not None else out_topk_ids, + *(out_logical_ids.stride() if out_logical_ids is not None else out_topk_ids.stride()), + logical_to_physical_map, + logical_replica_count, + expert_counter, + sample_index, + group_num=group_num, + group_expert_num=group_expert_num, + total_expert_num=total_expert_num, + group_topk_num=group_topk_num, + IS_SIGMOID=use_grouped_topk and scoring_func == "sigmoid", + USE_GROUPED_TOPK=use_grouped_topk, + HAS_CORRECTION_BIAS=use_grouped_topk and correction_bias is not None, + RETURN_LOGICAL_IDS=return_logical_ids, + EXPERT_GROUP_NUM=expert_group_num, + EXPERT_GROUP_SIZE=expert_group_size, + TOPK_NUM=topk, + TOPK_BLOCK_SIZE=triton.next_power_of_2(topk), + RENORMALIZE=renormalize, + GROUP_SCORE_USED_TOPK_NUM=group_score_used_topk_num, + COUNTER_NUM_EXPERTS=expert_counter.shape[1], + MAP_SLOTS=logical_to_physical_map.shape[1], + RECORD_LOAD=record_load, + SINGLE_TOKEN=token_num == 1, + num_warps=num_warps, + num_stages=1, + ) + return out_topk_weights, out_topk_ids, out_logical_ids diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py index 87cda6ce16..d2f59de480 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py @@ -21,7 +21,6 @@ from lightllm.utils.sgl_utils import sgl_ops from typing import Callable, List, Optional, Tuple from lightllm.common.basemodel.triton_kernel.fused_moe.softmax_topk import softmax_topk -from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType def fused_topk( @@ -168,12 +167,4 @@ def select_experts( hidden_states=hidden_states, gating_output=router_logits, topk=top_k, renormalize=renormalize ) - ######################################## warning ################################################## - # here is used to match autotune feature, make topk_ids more random - if Autotuner.is_kernel_autotune_warmup(AutotuneKernelType.GENERAL): - rand_gen = torch.Generator(device="cuda") - rand_gen.manual_seed(router_logits.shape[0]) - router_logits = torch.randn(size=router_logits.shape, generator=rand_gen, dtype=torch.float32, device="cuda") - _, topk_ids = torch.topk(router_logits, k=top_k, dim=1) - return topk_weights, topk_ids diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index b721585596..d18c4a780f 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -201,7 +201,15 @@ def new_deepep_group( self.ll_num_tokens = prefill_num_max_dispatch_tokens_per_rank self.ll_decode_num_tokens = decode_num_max_dispatch_tokens_per_rank self.ll_hidden = hidden_size - self.ll_num_experts = n_routed_experts + total_redundant_experts = ( + get_env_start_args().eplb_num_redundant_experts_per_rank * global_world_size + if get_env_start_args().enable_prefill_eplb + else 0 + ) + self.ll_prefill_num_experts = n_routed_experts + total_redundant_experts + # EPLB's redundant rows are a prefill-only physical layout; decode + # always routes the logical expert space. + self.ll_decode_num_experts = n_routed_experts self.ep_buffer = deep_ep.ElasticBuffer( deepep_group, num_max_tokens_per_rank=self.ll_num_tokens, @@ -241,7 +249,10 @@ def new_deepep_group( # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 # 空闲的本地 RDMA storage 复用为分块 grouped GEMM 的临时 workspace。 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 + self.ll_decode_num_tokens, + self.ll_hidden, + global_world_size, + self.ll_decode_num_experts, ) microbatch_count = len(self.groups) min_prefill_reuse_buffer_bytes = _calculate_min_chunked_expanded_moe_reuse_buffer_bytes( @@ -262,7 +273,7 @@ def new_deepep_group( deepep_group, num_rdma_bytes=num_rdma_bytes, low_latency_mode=True, - num_qps_per_rank=(self.ll_num_experts // global_world_size), + num_qps_per_rank=(self.ll_decode_num_experts // global_world_size), ) if enable_mega_moe_buffer: @@ -275,22 +286,26 @@ def new_deepep_group( self.ep_mega_moe_buffer = deep_gemm.get_symm_buffer_for_mega_moe( deepep_group, - self.ll_num_experts, + self.ll_decode_num_experts, self.ll_num_tokens, num_experts_per_tok, self.ll_hidden, moe_intermediate_size, ) logger.info( - "Initialize DeepEP MoE buffers: low_latency=%s, mega_moe=%s, expert_quant_method_names=%s", + "Initialize DeepEP MoE buffers: low_latency=%s, mega_moe=%s, " + "ll_prefill_num_experts=%s, ll_decode_num_experts=%s, expert_quant_method_names=%s", enable_low_latency_buffer, enable_mega_moe_buffer, + self.ll_prefill_num_experts, + self.ll_decode_num_experts, 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) + theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_prefill_num_experts, num_experts_per_tok) + low_latency_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_decode_num_experts, num_experts_per_tok) + self._set_num_sms_for_deep_gemm(theoretical_sms, low_latency_sms) - def _set_num_sms_for_deep_gemm(self, deepep_sms: int): + def _set_num_sms_for_deep_gemm(self, deepep_sms: int, low_latency_sms: int): try: try: from deep_gemm.jit_kernels.utils import set_num_sms @@ -299,9 +314,12 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): device_sms = get_device_sm_count() deepep_sms = max(0, min(deepep_sms, max(device_sms - 2, 0))) + low_latency_sms = max(0, min(low_latency_sms, max(device_sms - 2, 0))) self.ep_num_sms = deepep_sms if self.ep_low_latency_buffer is not None: - deep_ep.Buffer.set_num_sms(deepep_sms - deepep_sms % 2) + # This setting controls the legacy low-latency buffer; keep + # its SM reservation based on decode's logical expert count. + deep_ep.Buffer.set_num_sms(low_latency_sms - low_latency_sms % 2) set_num_sms(max(device_sms - deepep_sms, 2)) except BaseException as e: logger.warning(f"set num sms for deep_gemm failed: {e}") @@ -343,7 +361,7 @@ def clear_deepep_buffer(self): """ if self.ep_low_latency_buffer is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( - self.ll_decode_num_tokens, self.ll_hidden, self.ll_num_experts + self.ll_decode_num_tokens, self.ll_hidden, self.ll_decode_num_experts ) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 9d90ff69ad..a2bb02879d 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -771,6 +771,18 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="""Disable the prefill expert balance monitor enabled by default for EP-MoE.""", ) + parser.add_argument( + "--enable_prefill_eplb", + action="store_true", + help="""Enable online expert load balancing for prefill only.""", + ) + parser.add_argument( + "--eplb_num_redundant_experts_per_rank", + type=int, + default=2, + help="""Number of redundant physical experts per EP rank for each MoE layer used by prefill EPLB. + The value must be greater than 0.""", + ) parser.add_argument( "--enable_fused_shared_experts", action="store_true", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 1c433ec60e..31e5f39dbc 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -25,6 +25,7 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args +from lightllm.utils.device_utils import is_sm100_gpu logger = init_logger(__name__) @@ -158,6 +159,16 @@ def _launch_subprocesses(args: StartArgs): if args.enable_dp_prefill_balance: assert args.enable_tpsp_mix_mode and args.dp > 1, "need set --enable_tpsp_mix_mode firstly and --dp > 1" + if args.enable_prefill_eplb: + assert args.enable_ep_moe, "--enable_prefill_eplb requires --enable_ep_moe" + assert not args.enable_prefill_cudagraph, "--enable_prefill_eplb does not support --enable_prefill_cudagraph" + # EPLB updates expert weights in place, but SM100 Mega-MoE caches transformed weights by tensor data_ptr. + assert not is_sm100_gpu(), "--enable_prefill_eplb does not support SM100" + assert ( + args.eplb_num_redundant_experts_per_rank > 0 + ), "--eplb_num_redundant_experts_per_rank must be greater than 0" + assert args.mtp_mode is None, "--enable_prefill_eplb does not support MTP modes" + if args.enable_ep_moe: allowed_ep_prefill_att_backends = {"auto", "fa3", "triton", "flashqla"} for backend in args.llm_prefill_att_backend: diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 6ad7a9b9d8..0dd4881831 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -187,6 +187,8 @@ class StartArgs: ) enable_ep_moe: bool = field(default=False) disable_ep_balance_monitor: bool = field(default=False) + enable_prefill_eplb: bool = field(default=False) + eplb_num_redundant_experts_per_rank: int = field(default=2) enable_fused_shared_experts: bool = field(default=False) mtp_mode: Optional[str] = field( default=None, diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index c4317df652..3291a17d0f 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -256,14 +256,22 @@ def init_model(self, kvargs): prof_name = f"lightllm-model_backend-node{self.node_rank}_dev{get_current_device_id()}" prof_mode = self.args.enable_profiling self.profiler = ProcessProfiler(mode=prof_mode, name=prof_name, use_multi_thread=True) if prof_mode else None + if self.args.enable_prefill_eplb: + from lightllm.server.router.model_infer.mode_backend.eplb_manager import EPLBManager + self.model.eplb_manager = EPLBManager(self.model) + dist.barrier() + + self.start_infer_loops() + return + + def start_infer_loops(self): # 启动infer_loop_thread, 启动两个线程进行推理,对于具备双batch推理折叠得场景 # 可以降低 cpu overhead,大幅提升gpu得使用率。 self.infer_loop_thread = threading.Thread(target=self.infer_loop, daemon=True) self.infer_loop_thread.start() self.infer_loop_thread1 = threading.Thread(target=self.infer_loop, daemon=True) self.infer_loop_thread1.start() - return def init_custom(self): pass diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 2326c8c515..12b5195551 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -61,6 +61,11 @@ def infer_loop(self): event_pack.wait_to_forward() + # EPLB polling performs collectives and may commit weights, so keep all ranks ordered before normal + # collectives/forward. + if self.model.eplb_manager is not None: + self.model.eplb_manager.poll() + self._try_read_new_reqs() prefill_reqs, decode_reqs = self._get_classed_reqs( diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index ce21d01987..3070645b7e 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -122,6 +122,11 @@ def infer_loop(self): event_pack.wait_to_forward() + # EPLB polling performs collectives and may commit weights, so keep all ranks ordered before normal + # collectives/forward. + if self.model.eplb_manager is not None: + self.model.eplb_manager.poll() + self._try_read_new_reqs() prefill_reqs, decode_reqs = self._get_classed_reqs( diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py new file mode 100644 index 0000000000..aa597e6a11 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -0,0 +1,558 @@ +import threading +import time +from typing import Dict, Optional + +import torch +import torch.distributed as dist + +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_logical_to_physical_maps_for_layers, + plan_redundant_experts, + select_improving_placements, +) +from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + NixlEPLBTransfer, + align_target_placement, + build_transfer_plan, +) +from lightllm.utils.dist_utils import get_global_rank, get_global_world_size, get_node_world_size +from lightllm.utils.envs_utils import ( + get_eplb_placement_stickiness, + get_eplb_rebalance_gain_threshold, + get_prefill_eplb_step_interval, +) +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) +EPLB_MIN_AVG_TOKENS_PER_EXPERT = 100 +EPLB_EXPERT_ALIGNMENT = 128 +EPLB_CONTROL_ERROR = -1 +EPLB_STEADY_SAMPLE_STEPS = 4 + + +class EPLBManager: + """Online EPLB with asynchronous GPU expert migration.""" + + def __init__(self, model: TpPartBaseModel): + self.weights = _find_fused_moe_weights(model) + assert self.weights, "EPLB requires at least one EP MoE layer" + self.global_rank = get_global_rank() + self.world_size = get_global_world_size() + self.node_world_size = get_node_world_size() + self._eplb_states = [weight.expert_parallel_state.eplb for weight in self.weights] + self.step_interval = get_prefill_eplb_step_interval() + self.rebalance_gain_threshold = get_eplb_rebalance_gain_threshold() + self.placement_stickiness = get_eplb_placement_stickiness() + self.sampling_interval = self.step_interval + self.prefill_steps = 0 + routed = {weight.expert_parallel_state.num_logical_experts for weight in self.weights} + redundant = {state.num_redundant_experts_per_rank for state in self._eplb_states} + assert len(routed) == len(redundant) == 1 + self.num_logical_experts = routed.pop() + self.num_redundant_experts_per_rank = redundant.pop() + self.current_placement = torch.stack( + [state.initial_redundant_expert_ids_by_rank for state in self._eplb_states] + ) + self.in_flight = False + self.target_placement = None + self.target_metadata = None + self.in_flight_started_at = None + self.evaluation_in_flight = False + self._evaluation_lock = threading.Lock() + self._evaluation_result = None + self._evaluation_error = None + self._evaluation_thread = None + # A fresh manager starts with one continuous base window. After a + # sufficient evaluation, steady state returns to the cheap sparse + # probe. An insufficient sparse probe schedules one fresh continuous + # base window before the next fixed sampling boundary. + self._continuous_collection_start_step: Optional[int] = None + self._continuous_collection_end_step: Optional[int] = self.step_interval + self._sampling_pending = False + self._steady_collection_end_step: Optional[int] = None + self._reset_recorded_samples() + self._set_recording(True) + # Keep background evaluation collectives separate from the main-thread + # control/poll collectives: their ordering is intentionally independent. + self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") + self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") + self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") + # This control-group scalar is only touched from the main inference + # thread, never by the background evaluation thread. + self._control_ready_count = torch.empty(1, dtype=torch.int32) + self.transfer = NixlEPLBTransfer(self.weights, self.transfer_group, self.global_rank, self.world_size) + if self.global_rank == 0: + logger.info( + "eplb enabled " + f"layers={len(self.weights)} num_logical_experts={self.num_logical_experts} " + f"num_redundant_experts_per_rank={self.num_redundant_experts_per_rank} " + f"step_interval={self.step_interval} " + f"rebalance_gain_threshold={self.rebalance_gain_threshold:.4f} " + f"placement_stickiness={self.placement_stickiness:.4f}" + ) + + def poll(self): + """Poll only from a globally ordered pre-forward boundary.""" + if self.in_flight: + self._poll_in_flight() + return + if self.evaluation_in_flight and self._evaluation_ready_on_all_ranks(): + self._poll_evaluation() + + def step(self): + if self.in_flight or self.evaluation_in_flight: + return + self.prefill_steps += 1 + continuous_start = self._continuous_collection_start_step + continuous_end = self._continuous_collection_end_step + if continuous_end is not None: + # 启动或稀疏样本不足时连续采样,保证负载统计可靠。 + if continuous_start is not None and self.prefill_steps == continuous_start: + self._set_recording(True) + if self.prefill_steps >= continuous_end: + self._start_evaluation() + return + sampling_interval = self.sampling_interval + phase = self.prefill_steps % sampling_interval + # _sampling_pending=True 表示稳态采样窗口已启动,防止重复启动;到期后清除标记并评估。 + if self._sampling_pending: + steady_collection_end_step = self._steady_collection_end_step + if steady_collection_end_step is None or self.prefill_steps >= steady_collection_end_step: + self._clear_steady_collection() + self._start_evaluation() + return + if sampling_interval == 1: + self._start_evaluation() + return + if phase == sampling_interval - self._steady_sample_window_steps(): + # 稳态仅在周期末采样少量 step,降低路由计数和评估开销。 + self._start_steady_sampling_window(self.prefill_steps + self._steady_sample_window_steps()) + + def _set_recording(self, enabled: bool): + for state in self._eplb_states: + state.recording = enabled + + def _reset_recorded_samples(self): + counters = [state.route_counter for state in self._eplb_states] + if counters: + torch._foreach_zero_(counters) + for state in self._eplb_states: + state.recorded_sample_count = 0 + + def _control_count(self, value: int) -> torch.Tensor: + """Return the main-thread-only reusable control collective scalar.""" + return self._control_ready_count.fill_(value) + + def _clear_continuous_collection(self): + self._continuous_collection_start_step = None + self._continuous_collection_end_step = None + + def _clear_steady_collection(self): + self._sampling_pending = False + self._steady_collection_end_step = None + + def _steady_sample_window_steps(self) -> int: + return min(EPLB_STEADY_SAMPLE_STEPS, self.sampling_interval) + + def _start_steady_sampling_window(self, collection_end_step: int): + """Start the fixed sparse window without moving its evaluation boundary.""" + self._reset_recorded_samples() + self._sampling_pending = True + self._steady_collection_end_step = collection_end_step + self._set_recording(True) + + def _begin_continuous_collection(self): + minimum_end = self.prefill_steps + self.step_interval + collection_end = -(-minimum_end // self.sampling_interval) * self.sampling_interval + self._reset_recorded_samples() + self._clear_steady_collection() + self._continuous_collection_start_step = collection_end - self.step_interval + self._continuous_collection_end_step = collection_end + self._set_recording(self._continuous_collection_start_step == self.prefill_steps) + + def _prepare_next_sampling_window(self): + """Clear the current window and arm the next sparse sampling window.""" + self._clear_continuous_collection() + self._clear_steady_collection() + if self.sampling_interval == 1: + self._reset_recorded_samples() + self._set_recording(True) + elif self.sampling_interval <= EPLB_STEADY_SAMPLE_STEPS: + # There is no later pre-boundary manager step at which to arm a + # full clamped window, so arm immediately but keep the same next + # fixed boundary. + self._start_steady_sampling_window(self.prefill_steps + self.sampling_interval) + else: + self._reset_recorded_samples() + self._set_recording(False) + + @staticmethod + def _recent_ring_samples(counter: torch.Tensor, recorded_sample_count: int) -> torch.Tensor: + """Return the newest ring rows in chronological order.""" + capacity = counter.shape[0] + available = min(recorded_sample_count, capacity) + if available == 0: + return counter[:0] + start = (recorded_sample_count - available) % capacity + indices = (torch.arange(available, dtype=torch.int64, device=counter.device) + start) % capacity + return counter.index_select(0, indices) + + def _collect_local_samples(self) -> torch.Tensor: + counters = [state.route_counter for state in self._eplb_states] + capacities = [counter.shape[0] for counter in counters] + if len(set(capacities)) != 1 or any(counter.ndim != 2 for counter in counters): + raise RuntimeError("EPLB sample capacities differ between layers") + counts = [state.recorded_sample_count for state in self._eplb_states] + if len(set(counts)) != 1: + raise RuntimeError("EPLB recorded sample counts differ between layers") + sample_count = counts[0] + # Validate the metadata before copying the newest rows to the CPU. + metadata = torch.tensor([sample_count, -sample_count, capacities[0], -capacities[0]], dtype=torch.int64) + dist.all_reduce(metadata, op=dist.ReduceOp.MIN, group=self.evaluation_group) + if metadata[0] != -metadata[1] or metadata[2] != -metadata[3]: + raise RuntimeError("EPLB recorded sample count or capacity differs between ranks") + # Stack the fixed-size ring buffers in one GPU launch. Slicing each + # layer before stacking turns a single launch into one index_select per + # MoE layer and is measurably slower in the normal sparse path. + counter_samples = torch.stack(counters, dim=1) + return self._recent_ring_samples(counter_samples, sample_count).cpu() + + def _commit_layer_metadata(self, layer_index: int): + eplb_state = self._eplb_states[layer_index] + logical_to_physical, replica_count = self.target_metadata[layer_index] + eplb_state.logical_to_physical_map.copy_(logical_to_physical, non_blocking=True) + eplb_state.logical_replica_count.copy_(replica_count, non_blocking=True) + + def _finish_rebalance(self): + self.current_placement = self.target_placement + self.target_placement = None + self.target_metadata = None + self.in_flight = False + self._prepare_next_sampling_window() + if self.global_rank == 0: + logger.info(f"eplb completed wall_time={time.time() - self.in_flight_started_at:.2f}s") + + def _poll_in_flight(self): + local_error = None + try: + pending = self.transfer.pending_layers() + except BaseException as exc: + pending = [] + local_error = exc + ready_count = self._control_count(EPLB_CONTROL_ERROR if local_error is not None else len(pending)) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) + ready_count = int(ready_count.item()) + if ready_count < 0: + if local_error is not None: + raise RuntimeError("EPLB transfer worker failed on this rank") from local_error + raise RuntimeError("EPLB transfer worker failed on another rank") + if ready_count == 0: + return + if ready_count > len(pending) or ready_count > len(self.in_flight_layers): + raise RuntimeError("EPLB global ready count exceeds the local ordered prefix") + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + # Previous forward is queued on the shared overlap stream; order the + # live-weight commit after it. The subsequent wait orders the next forward. + torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) + for layer_index, buffer_index in pending[:ready_count]: + if layer_index != self.in_flight_layers[0]: + raise RuntimeError( + f"EPLB pending layer {layer_index} does not match expected {self.in_flight_layers[0]}" + ) + self.transfer.commit(layer_index, buffer_index, lambda: self._commit_layer_metadata(layer_index)) + self.in_flight_layers.pop(0) + if not self.in_flight_layers: + self.transfer.finish() + self._finish_rebalance() + + def _plan_and_broadcast(self, global_load: torch.Tensor): + """Plan on rank zero and share the serializable result on the evaluation group.""" + result = None + local_error = None + if self.global_rank == 0: + try: + minimum = self.num_logical_experts * EPLB_MIN_AVG_TOKENS_PER_EXPERT + layer_samples = global_load.sum(dim=(0, 2, 3)) + if torch.any(layer_samples < minimum): + result = { + "kind": "insufficient", + "minimum_layer_samples": int(layer_samples.min().item()), + "minimum": minimum, + } + else: + candidate = plan_redundant_experts( + global_load, + self.world_size, + self.num_redundant_experts_per_rank, + expert_alignment=EPLB_EXPERT_ALIGNMENT, + node_world_size=self.node_world_size, + current_placement=self.current_placement, + stickiness=self.placement_stickiness, + ) + placement, improved, metrics, before_load, after_load = select_improving_placements( + global_load, + self.current_placement, + candidate, + expert_alignment=EPLB_EXPERT_ALIGNMENT, + node_world_size=self.node_world_size, + rebalance_gain_threshold=self.rebalance_gain_threshold, + ) + if bool(torch.any(improved)): + # A planner placement identifies experts by rank, not + # by redundant slot. Canonicalize every selected row + # before broadcasting so transfer, metadata, and the + # next current_placement all describe the same live + # physical expert rows. + placement = placement.clone() + for layer_index in torch.nonzero(improved, as_tuple=False).flatten().tolist(): + placement[layer_index] = align_target_placement( + self.current_placement[layer_index], placement[layer_index] + ) + result = { + "kind": "planned" if bool(torch.any(improved)) else "no_improvement", + "placement": placement, + "improved": improved, + "before": _imbalance_summary(before_load), + "after": _imbalance_summary(after_load), + **metrics, + } + except BaseException as exc: + local_error = exc + result = {"kind": "error", "message": f"{type(exc).__name__}: {exc}"} + if self.world_size > 1: + result_list = [result] + dist.broadcast_object_list(result_list, src=0, group=self.evaluation_group) + result = result_list[0] + if result["kind"] == "error": + if local_error is not None: + raise RuntimeError("EPLB planner failed on rank zero") from local_error + raise RuntimeError(f"EPLB planner failed on rank zero: {result['message']}") + return result + + def _evaluate_after_event(self, event: torch.cuda.Event): + """Run the CPU/Gloo planning phase after the frozen CUDA counters are ready.""" + try: + torch.cuda.set_device(self._eplb_states[0].route_counter.device) + event.synchronize() + local_load = self._collect_local_samples() + recorded_sample_count = int(local_load.shape[0]) + sample_window_steps = ( + self.step_interval + if self._continuous_collection_end_step is not None + else self._steady_sample_window_steps() + ) + num_nodes = self.world_size // self.node_world_size + # Preserve source nodes until physical-replica loads are combined; + # DeepEP applies expert alignment after traffic from all sources + # reaches each destination expert. + global_load = torch.zeros((*local_load.shape[:2], num_nodes, local_load.shape[2]), dtype=local_load.dtype) + global_load[:, :, self.global_rank // self.node_world_size] = local_load + dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) + result = self._plan_and_broadcast(global_load) + result["recorded_sample_count"] = recorded_sample_count + result["sample_window_steps"] = sample_window_steps + if result["kind"] == "planned": + metadata = [None] * len(self.weights) + layer_plans = [] + improved_layer_indices = torch.nonzero(result["improved"], as_tuple=False).flatten() + if improved_layer_indices.numel(): + maps_for_improved_layers, counts_for_improved_layers = build_logical_to_physical_maps_for_layers( + result["placement"][improved_layer_indices], + self.num_logical_experts, + source_rank=self.global_rank, + node_world_size=self.node_world_size, + ) + for improved_layer_offset, layer_index in enumerate(improved_layer_indices.tolist()): + placement = result["placement"][layer_index] + metadata[layer_index] = ( + maps_for_improved_layers[improved_layer_offset], + counts_for_improved_layers[improved_layer_offset], + ) + layer_plans.append( + ( + layer_index, + build_transfer_plan( + self.current_placement[layer_index], + placement, + self.num_logical_experts, + self.world_size, + self.node_world_size, + ), + ) + ) + result["metadata"] = metadata + result["layer_plans"] = layer_plans + with self._evaluation_lock: + self._evaluation_result = result + except BaseException as exc: + with self._evaluation_lock: + self._evaluation_error = exc + + def _start_evaluation(self): + with self._evaluation_lock: + self._evaluation_result = None + self._evaluation_error = None + self._set_recording(False) + event = torch.cuda.Event() + event.record(torch.cuda.current_stream()) + self.evaluation_in_flight = True + self._evaluation_thread = threading.Thread(target=self._evaluate_after_event, args=(event,), daemon=True) + self._evaluation_thread.start() + + def _poll_evaluation(self): + if not self.evaluation_in_flight: + return False + with self._evaluation_lock: + error = self._evaluation_error + result = self._evaluation_result + if error is not None or result is not None: + self._evaluation_result = None + self._evaluation_error = None + if error is not None: + self._evaluation_thread.join() + self.evaluation_in_flight = False + self._evaluation_thread = None + raise error + if result is None: + return True + self._evaluation_thread.join() + self.evaluation_in_flight = False + self._evaluation_thread = None + if result["kind"] == "insufficient": + from_continuous_window = self._continuous_collection_end_step is not None + if from_continuous_window: + self.sampling_interval = min(self.sampling_interval * 4, self.step_interval * 16) + self._prepare_next_sampling_window() + else: + self._begin_continuous_collection() + if self.global_rank == 0: + if from_continuous_window: + logger.info( + "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " + "next_sampling_interval=%s recorded_sample_count=%s sample_window_steps=%s", + self.prefill_steps, + result["minimum_layer_samples"], + result["minimum"], + self.sampling_interval, + result.get("recorded_sample_count"), + result.get("sample_window_steps"), + ) + else: + logger.info( + "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " + "scheduled_fresh_window_start=%s scheduled_fresh_window_end=%s " + "recorded_sample_count=%s sample_window_steps=%s", + self.prefill_steps, + result["minimum_layer_samples"], + result["minimum"], + self._continuous_collection_start_step, + self._continuous_collection_end_step, + result.get("recorded_sample_count"), + result.get("sample_window_steps"), + ) + return False + if result["kind"] == "no_improvement": + self.sampling_interval = min(self.sampling_interval * 4, self.step_interval * 16) + if self.global_rank == 0: + logger.info( + "eplb skip rearrangement: no model improvement model_imbalance_ratio=%.4f " + "candidate_model_imbalance_ratio=%.4f candidate_rebalance_gain=%.4f " + "candidate_changed_layer_count=%s actual_changed_layer_count=0 next_sampling_interval=%s " + "recorded_sample_count=%s sample_window_steps=%s", + result["model_imbalance_ratio"], + result["candidate_model_imbalance_ratio"], + result["candidate_rebalance_gain"], + result["candidate_changed_layer_count"], + self.sampling_interval, + result.get("recorded_sample_count"), + result.get("sample_window_steps"), + ) + self._prepare_next_sampling_window() + return False + self._start_rebalance(result) + return True + + def _evaluation_ready_on_all_ranks(self) -> bool: + with self._evaluation_lock: + local_error = self._evaluation_error + local_result = self._evaluation_result + local_status = EPLB_CONTROL_ERROR if local_error is not None else int(local_result is not None) + ready_count = self._control_count(local_status) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) + ready_count = int(ready_count.item()) + if ready_count < 0: + if local_error is not None: + raise RuntimeError("EPLB evaluation failed on this rank") from local_error + raise RuntimeError("EPLB evaluation failed on another rank") + return bool(ready_count) + + def _start_rebalance(self, result): + placement = result["placement"] + layer_plans = result["layer_plans"] + self.sampling_interval = self.step_interval + self._clear_continuous_collection() + self._reset_recorded_samples() + self.target_placement = placement + self.target_metadata = result["metadata"] + self.in_flight_layers = [layer_index for layer_index, _ in layer_plans] + self.in_flight = True + self.in_flight_started_at = time.time() + self.transfer.start(layer_plans) + if self.global_rank == 0: + actual_changed_slot_count = sum(len(plan) for _, plan in layer_plans) + cross_node_transfer_count = sum( + step.src_rank // self.node_world_size != step.dst_rank // self.node_world_size + for _, plan in layer_plans + for step in plan + ) + logger.info( + "eplb started prefill_steps=%s max_before=%.4f max_after=%.4f p95_before=%.4f p95_after=%.4f " + "model_imbalance_ratio=%.4f candidate_model_imbalance_ratio=%.4f " + "candidate_rebalance_gain=%.4f candidate_changed_layer_count=%s " + "actual_changed_layer_count=%s actual_changed_slot_count=%s cross_node_transfer_count=%s " + "recorded_sample_count=%s sample_window_steps=%s", + self.prefill_steps, + result["before"]["max"], + result["after"]["max"], + result["before"]["p95"], + result["after"]["p95"], + result["model_imbalance_ratio"], + result["candidate_model_imbalance_ratio"], + result["candidate_rebalance_gain"], + result["candidate_changed_layer_count"], + len(layer_plans), + actual_changed_slot_count, + cross_node_transfer_count, + result.get("recorded_sample_count"), + result.get("sample_window_steps"), + ) + + +def _imbalance_summary(rank_load: torch.Tensor) -> Dict[str, float]: + if rank_load.ndim == 2: + critical = rank_load.max(dim=1).values + mean = rank_load.mean(dim=1) + elif rank_load.ndim == 3: + critical = rank_load.max(dim=2).values.sum(dim=0) + mean = rank_load.mean(dim=2).sum(dim=0) + else: + raise ValueError("rank_load must be [layers, ranks] or [samples, layers, ranks]") + layer_imbalance = critical / mean.clamp_min(1.0) + sorted_imbalance = torch.sort(layer_imbalance).values + p95_index = max(0, (95 * layer_imbalance.numel() + 99) // 100 - 1) + return { + "max": float(layer_imbalance.max().item()), + "p95": float(sorted_imbalance[p95_index].item()), + } + + +def _find_fused_moe_weights(model): + weights_by_id = {} + for layer in model.trans_layers_weight: + for value in getattr(layer, "__dict__", {}).values(): + if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: + weights_by_id[id(value)] = value + return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py new file mode 100644 index 0000000000..f0e00df81d --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -0,0 +1,655 @@ +"""Asynchronous expert-row migration for EPLB.""" +import os +import socket +import threading +from collections import defaultdict, deque +from dataclasses import dataclass +from typing import Dict, List, Sequence, Tuple + +import torch +import torch.distributed as dist + +from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_kernels import ( + eplb_push_copy, +) + + +@dataclass(frozen=True) +class TransferStep: + dst_rank: int + dst_slot: int + src_rank: int + src_local_row: int + + +def extract_expert_tensors(weight) -> List[Tuple[str, torch.Tensor]]: + result = [] + for pack_name in ("w13", "w2"): + pack = getattr(weight, pack_name) + for value_name in ("weight", "weight_scale", "weight_zero_point"): + tensor = getattr(pack, value_name, None) + if tensor is not None: + assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" + result.append((f"{pack_name}.{value_name}", tensor)) + return result + + +def commit_staging_rows( + live: torch.Tensor, + staging: torch.Tensor, + num_experts_per_rank: int, + changed_dst_slots: Sequence[int], +) -> None: + slots = sorted(set(changed_dst_slots)) + if not slots: + return + run_start = previous = slots[0] + for dst_slot in (*slots[1:], None): + if dst_slot is not None and dst_slot == previous + 1: + previous = dst_slot + continue + run_length = previous - run_start + 1 + live.narrow(0, num_experts_per_rank + run_start, run_length).copy_( + staging.narrow(0, run_start, run_length), non_blocking=True + ) + if dst_slot is not None: + run_start = previous = dst_slot + + +def align_target_placement(current: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + """Canonicalize a target row layout without moving retained experts. + + EPLB placement is rank-based: redundant slots on one rank are + interchangeable. Retained experts therefore keep their live physical + slot, while new experts fill freed slots in the planner's target-row + order. The returned placement is the single canonical layout that must + be used both for transfers and for published routing metadata. + """ + assert current.ndim == target.ndim == 2 + assert tuple(current.shape) == tuple(target.shape) + + current_rows = current.tolist() + target_rows = target.tolist() + aligned_target_rows = [] + for current_row, target_row in zip(current_rows, target_rows): + current_slots = {expert: slot for slot, expert in enumerate(current_row)} + target_experts = set(target_row) + aligned_row = list(current_row) + freed_slots = [slot for slot, expert in enumerate(current_row) if expert not in target_experts] + new_experts = [expert for expert in target_row if expert not in current_slots] + assert len(freed_slots) == len(new_experts) + for slot, expert in zip(freed_slots, new_experts): + aligned_row[slot] = expert + aligned_target_rows.append(aligned_row) + return target.new_tensor(aligned_target_rows) + + +def build_transfer_plan( + current: torch.Tensor, + target: torch.Tensor, + num_logical_experts: int, + world_size: int, + node_world_size: int, +) -> List[TransferStep]: + assert tuple(current.shape) == tuple(target.shape) == (world_size, current.shape[1]) + num_experts_per_rank = num_logical_experts // world_size + current_rows = current.tolist() + aligned_target_rows = align_target_placement(current, target).tolist() + # A logical expert has one primary row and at most one redundant row per + # rank, so this source list is already unique. Build it once instead of + # allocating/sorting a set for every destination slot. + candidates_by_expert = [ + [ + ( + expert // num_experts_per_rank, + expert % num_experts_per_rank, + ) + ] + for expert in range(num_logical_experts) + ] + for rank, row in enumerate(current_rows): + for slot, expert in enumerate(row): + candidates_by_expert[expert].append((rank, num_experts_per_rank + slot)) + source_load = [0] * world_size + plan = [] + for dst_rank in range(world_size): + for dst_slot, expert in enumerate(aligned_target_rows[dst_rank]): + if expert == current_rows[dst_rank][dst_slot]: + continue + src_rank, src_row = min( + candidates_by_expert[expert], + key=lambda item: ( + item[0] // node_world_size != dst_rank // node_world_size, + source_load[item[0]], + item[0], + item[1], + ), + ) + source_load[src_rank] += 1 + plan.append(TransferStep(dst_rank, dst_slot, src_rank, src_row)) + return plan + + +class _EPLBTransferBase: + """Shared live/staging buffers and publish/commit lifecycle.""" + + staging_depth = 1 + + def __init__(self, weights, transfer_group, global_rank, world_size): + self._eplb_states = [weight.expert_parallel_state.eplb for weight in weights] + self.transfer_group = transfer_group + self.global_rank = global_rank + self.world_size = world_size + self.num_experts_per_rank = weights[0].expert_parallel_state.num_primary_experts_per_rank + self.device = weights[0].w13.weight.device + self.live = [extract_expert_tensors(weight) for weight in weights] + self._validate_live_layout(weights) + num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank + self.staging = [ + [ + ( + name, + torch.empty( + (num_redundant_slots_per_rank,) + tuple(tensor.shape[1:]), + dtype=tensor.dtype, + device=tensor.device, + ), + ) + for name, tensor in self.live[0] + ] + for _ in range(self.staging_depth) + ] + self._release = [threading.Event() for _ in range(self.staging_depth)] + for release in self._release: + release.set() + self._error = None + self._consumed_events = [torch.cuda.Event() for _ in range(self.staging_depth)] + self._consumed_recorded = [False] * self.staging_depth + self._changed_dst_slots = [()] * self.staging_depth + self._pending = deque() + self._pending_lock = threading.Lock() + self._thread = None + self._needs_staging_reuse_barrier = False + + def _validate_live_layout(self, weights) -> None: + reference = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in self.live[0]] + num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank + for layer_index, (state, tensors) in enumerate(zip(self._eplb_states, self.live)): + layout = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in tensors] + assert layout == reference, f"EPLB layer {layer_index} has incompatible expert tensor layout" + assert ( + state.num_redundant_experts_per_rank == num_redundant_slots_per_rank + ), "EPLB redundant slot count must match" + + def _copy_layer(self, layer_index: int, plan: Sequence[TransferStep], staging) -> None: + raise NotImplementedError + + def _copy_batch(self, batch) -> None: + for layer_index, plan, _, staging in batch: + self._copy_layer(layer_index, plan, staging) + + def _start_transfer_generation(self) -> None: + """Prepare backend state after the in-flight worker check succeeds.""" + + def _finish_transfer_generation(self) -> None: + """Release backend state only after the migration worker has joined.""" + + def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]) -> None: + if self._thread is not None and self._thread.is_alive(): + raise RuntimeError("EPLB transfer is already in flight") + self._start_transfer_generation() + self._error = None + with self._pending_lock: + self._pending.clear() + + def worker() -> None: + try: + torch.cuda.set_device(self.device) + if not layer_plans: + self._finish_transfer_generation() + for batch_start in range(0, len(layer_plans), self.staging_depth): + batch = [] + for plan_index in range(batch_start, min(batch_start + self.staging_depth, len(layer_plans))): + layer_index, plan = layer_plans[plan_index] + buffer_index = plan_index % self.staging_depth + release = self._release[buffer_index] + # A buffer cannot be reused until its prior committed rows are no longer read by CUDA. + release.wait() + release.clear() + if self._consumed_recorded[buffer_index]: + self._consumed_events[buffer_index].synchronize() + self._changed_dst_slots[buffer_index] = tuple( + step.dst_slot for step in plan if step.dst_rank == self.global_rank + ) + batch.append((layer_index, plan, buffer_index, self.staging[buffer_index])) + if batch_start > 0 and self._needs_staging_reuse_barrier: + # All destinations must finish consuming the prior IPC staging generation + # before a source can reuse the peer buffer for this batch. + dist.barrier(group=self.transfer_group) + self._copy_batch(batch) + if batch_start + self.staging_depth >= len(layer_plans): + self._finish_transfer_generation() + with self._pending_lock: + self._pending.extend((layer_index, buffer_index) for layer_index, _, buffer_index, _ in batch) + except BaseException as exc: + self._error = exc + + self._thread = threading.Thread(target=worker, name=f"eplb-{self.backend}", daemon=True) + self._thread.start() + + def pending_layers(self): + if self._error is not None: + raise RuntimeError("EPLB migration worker failed") from self._error + with self._pending_lock: + return list(self._pending) + + def commit(self, layer_index: int, buffer_index: int, post_copy=None) -> None: + with self._pending_lock: + if not self._pending or self._pending[0] != (layer_index, buffer_index): + raise RuntimeError("EPLB commit does not match the pending FIFO") + self._pending.popleft() + changed_dst_slots = self._changed_dst_slots[buffer_index] + for (_, live), (_, staging) in zip(self.live[layer_index], self.staging[buffer_index]): + commit_staging_rows( + live, + staging, + self.num_experts_per_rank, + changed_dst_slots, + ) + if post_copy is not None: + post_copy() + self._consumed_events[buffer_index].record(torch.cuda.current_stream()) + self._consumed_recorded[buffer_index] = True + self._release[buffer_index].set() + + def finish(self) -> None: + """Wait for the released migration worker to exit before another rebalance.""" + thread = self._thread + if thread is None: + return + thread.join() + self._thread = None + if self._error is not None: + raise RuntimeError("EPLB migration worker failed") from self._error + + +class NixlEPLBTransfer(_EPLBTransferBase): + """GPU-direct UCX/NIXL EPLB transfer. Initialization errors are fatal.""" + + backend = "nixl" + staging_depth = 8 + _DEFAULT_UCX_TLS = "self,sm,cuda_ipc,cuda_copy,rc_x" + + def __init__(self, weights, transfer_group, global_rank, world_size): + super().__init__(weights, transfer_group, global_rank, world_size) + self._nixl_agent = None + self._registered_descs = None + self._remote_agents: Dict[int, str] = {} + self._remote_layouts = {} + self._xfer_cache = {} + self._used_xfer_cache_keys = set() + self._ipc_staging = {} + self._same_node_ranks = set() + self._cross_node_ranks = set() + self._push_stream = torch.cuda.Stream(device=self.device) + self._push_descriptor_cache = {} + self._used_push_descriptor_cache_keys = set() + try: + self._init_ipc_metadata() + if self._cross_node_ranks: + os.environ.setdefault("UCX_TLS", self._DEFAULT_UCX_TLS) + try: + import nixl + except Exception as exc: + raise RuntimeError("NIXL EPLB backend requires the nixl package for cross-node transfer") from exc + agent_name = f"lightllm-eplb-{socket.gethostname()}-{os.getpid()}-rank-{global_rank}" + config = nixl.nixl_agent_config(enable_prog_thread=True, enable_listen_thread=False, backends=["UCX"]) + self._nixl_agent = nixl.nixl_agent(agent_name, config) + reg_tensors = [tensor for layer in self.live for _, tensor in layer] + [ + tensor for staging in self.staging for _, tensor in staging + ] + self._registered_descs = self._nixl_agent.get_reg_descs(reg_tensors) + self._nixl_agent.register_memory(self._registered_descs, backends=["UCX"]) + self._init_remote_metadata() + except Exception as exc: + self.shutdown() + if isinstance(exc, RuntimeError): + raise + raise RuntimeError("NIXL EPLB initialization failed") from exc + + def _local_layout(self): + return [ + [(name, tensor.data_ptr(), tensor.get_device(), tensor[0].nbytes) for name, tensor in layer] + for layer in self.live + ] + + def _init_ipc_metadata(self) -> None: + hostnames = [None] * self.world_size + dist.all_gather_object(hostnames, socket.gethostname(), group=self.transfer_group) + local_hostname = hostnames[self.global_rank] + self._needs_staging_reuse_barrier = len(set(hostnames)) < len(hostnames) + self._same_node_ranks = {rank for rank, hostname in enumerate(hostnames) if hostname == local_hostname} + self._cross_node_ranks = set(range(self.world_size)) - self._same_node_ranks + for layer in self.live: + for name, tensor in layer: + if name.endswith(".weight") and tensor[0].nbytes % 16: + raise RuntimeError(f"NIXL source-push requires 16-byte aligned weight rows: {name}") + + from lightllm.server.router.model_infer.mode_backend.pd.p2p_fix import ( + p2p_fix_rebuild_cuda_tensor, + reduce_tensor, + ) + + exports = {} + for target_rank in self._same_node_ranks - {self.global_rank}: + exports[target_rank] = { + "staging": [ + [(name, tuple(tensor.shape), tensor.dtype, reduce_tensor(tensor)[1]) for name, tensor in staging] + for staging in self.staging + ], + } + all_exports = [None] * self.world_size + dist.all_gather_object(all_exports, exports, group=self.transfer_group) + + torch.cuda.set_device(self.device) + for dst_rank in self._same_node_ranks - {self.global_rank}: + metadata = all_exports[dst_rank].get(self.global_rank) + if metadata is None or len(metadata["staging"]) != self.staging_depth: + raise RuntimeError(f"NIXL IPC destination rank {dst_rank} has incompatible staging metadata") + rebuilt_staging = [] + for remote_staging, local_staging in zip(metadata["staging"], self.staging): + if len(remote_staging) != len(local_staging): + raise RuntimeError(f"NIXL IPC destination rank {dst_rank} staging tensor count mismatch") + rebuilt = [] + for (name, shape, dtype, args), (local_name, local_tensor) in zip(remote_staging, local_staging): + if name != local_name or shape != tuple(local_tensor.shape) or dtype != local_tensor.dtype: + raise RuntimeError(f"NIXL IPC destination rank {dst_rank} staging layout mismatch for {name}") + tensor = p2p_fix_rebuild_cuda_tensor(*args) + if tuple(tensor.shape) != shape or tensor.dtype != dtype or tensor.device != local_tensor.device: + raise RuntimeError( + f"NIXL IPC destination rank {dst_rank} staging rebuild validation failed for {name}" + ) + rebuilt.append((name, tensor)) + rebuilt_staging.append(rebuilt) + self._ipc_staging[dst_rank] = rebuilt_staging + + def _init_remote_metadata(self) -> None: + metadata = self._nixl_agent.get_agent_metadata() + all_metadata = [None] * self.world_size + all_layouts = [None] * self.world_size + dist.all_gather_object(all_metadata, metadata, group=self.transfer_group) + dist.all_gather_object(all_layouts, self._local_layout(), group=self.transfer_group) + for rank in self._cross_node_ranks: + layout = all_layouts[rank] + if len(layout) != len(self.live): + raise RuntimeError(f"NIXL remote rank {rank} has incompatible layer layout") + self._remote_agents[rank] = self._nixl_agent.add_remote_agent(all_metadata[rank]) + self._remote_layouts[rank] = layout + + def _wait_xfers(self, xfers) -> None: + pending = [] + for item in xfers: + state = self._nixl_agent.transfer(item[2]) + if state == "ERR": + raise RuntimeError("NIXL READ post failed") + if state == "PROC": + pending.append(item) + while pending: + remaining = [] + for item in pending: + state = self._nixl_agent.check_xfer_state(item[2]) + if state == "ERR": + raise RuntimeError("NIXL READ transfer failed") + if state != "DONE": + remaining.append(item) + pending = remaining + + def _release_xfers(self, xfers) -> None: + unreleased = [] + errors = [] + for local_dlist, remote_dlist, xfer in xfers: + remaining = [local_dlist, remote_dlist, xfer] + for remaining_index, handle, release in ( + (2, xfer, self._nixl_agent.release_xfer_handle), + (1, remote_dlist, self._nixl_agent.release_dlist_handle), + (0, local_dlist, self._nixl_agent.release_dlist_handle), + ): + if handle is not None: + try: + release(handle) + except Exception as exc: + errors.append(exc) + else: + remaining[remaining_index] = None + if any(handle is not None for handle in remaining): + unreleased.append(tuple(remaining)) + if errors: + error = RuntimeError("NIXL transfer handle release failed") + error.unreleased_xfers = unreleased + raise error from errors[0] + + @staticmethod + def _contiguous_runs(steps): + ordered = sorted(steps, key=lambda step: (step.src_local_row, step.dst_slot)) + runs = [] + for step in ordered: + if ( + runs + and step.src_local_row == runs[-1][-1].src_local_row + 1 + and step.dst_slot == runs[-1][-1].dst_slot + 1 + ): + runs[-1].append(step) + else: + runs.append([step]) + return runs + + @staticmethod + def _remote_read_cache_key(src_rank: int, entries): + return ( + src_rank, + tuple( + ( + layer_index, + tuple((step.src_local_row, step.dst_slot) for step in run), + tuple(tensor.data_ptr() for _, tensor in staging), + ) + for layer_index, run, staging in entries + ), + ) + + def _push_staging(self, dst_rank: int, buffer_index: int): + return self.staging[buffer_index] if dst_rank == self.global_rank else self._ipc_staging[dst_rank][buffer_index] + + def _cached_descriptor_tensors(self, copies): + key = tuple((source.data_ptr(), destination.data_ptr()) for destination, source in copies) + cached = self._push_descriptor_cache.get(key) + if cached is None: + src_ptrs = torch.tensor([source.data_ptr() for _, source in copies], dtype=torch.int64, device=self.device) + dst_ptrs = torch.tensor( + [destination.data_ptr() for destination, _ in copies], dtype=torch.int64, device=self.device + ) + cached = (src_ptrs, dst_ptrs) + self._push_descriptor_cache[key] = cached + self._used_push_descriptor_cache_keys.add(key) + return cached + + def _push_same_node(self, dst_rank: int, entries) -> None: + staging_by_buffer = {buffer_index: self._push_staging(dst_rank, buffer_index) for _, _, buffer_index in entries} + weight_groups = defaultdict(list) + small_copies = [] + for layer_index, run, buffer_index in entries: + staging = staging_by_buffer[buffer_index] + source_layer = self.live[layer_index] + first = run[0] + run_len = len(run) + for (name, source_tensor), (staging_name, staging_tensor) in zip(source_layer, staging): + if name != staging_name: + raise RuntimeError("NIXL source-push staging tensor name mismatch") + source_rows = source_tensor.narrow(0, first.src_local_row, run_len) + destination_rows = staging_tensor.narrow(0, first.dst_slot, run_len) + if name.endswith(".weight"): + if destination_rows.nbytes % 16: + raise RuntimeError(f"NIXL source-push requires 16-byte aligned weight rows: {name}") + weight_groups[destination_rows.nbytes].append((destination_rows, source_rows)) + else: + small_copies.append((destination_rows, source_rows)) + with torch.cuda.stream(self._push_stream): + for nbytes, copies in weight_groups.items(): + src_ptrs, dst_ptrs = self._cached_descriptor_tensors(copies) + eplb_push_copy(src_ptrs, dst_ptrs, nbytes) + for destination_rows, source_rows in small_copies: + destination_rows.copy_(source_rows, non_blocking=True) + + def _get_remote_read(self, src_rank: int, entries): + cache_key = self._remote_read_cache_key(src_rank, entries) + cached = self._xfer_cache.get(cache_key) + if cached is not None: + self._used_xfer_cache_keys.add(cache_key) + return cached + local_descs = [] + remote_descs = [] + local_dlist = remote_dlist = xfer = None + try: + for layer_index, run, staging in entries: + remote_layer = self._remote_layouts[src_rank][layer_index] + if len(remote_layer) != len(staging): + raise RuntimeError(f"NIXL remote rank {src_rank} has incompatible layer layout") + first = run[0] + run_len = len(run) + for tensor_index, (_, staging_tensor) in enumerate(staging): + name, remote_ptr, remote_device, remote_nbytes = remote_layer[tensor_index] + if ( + name != self.live[layer_index][tensor_index][0] + or remote_nbytes != staging_tensor[first.dst_slot].nbytes + ): + raise RuntimeError(f"NIXL remote rank {src_rank} descriptor range mismatch") + local_descs.append( + ( + staging_tensor[first.dst_slot].data_ptr(), + run_len * remote_nbytes, + staging_tensor.get_device(), + ) + ) + remote_descs.append( + (remote_ptr + first.src_local_row * remote_nbytes, run_len * remote_nbytes, remote_device) + ) + local_dlist = self._nixl_agent.prep_xfer_dlist( + "NIXL_INIT_AGENT", self._nixl_agent.get_xfer_descs(local_descs, "VRAM"), backends=["UCX"] + ) + remote_dlist = self._nixl_agent.prep_xfer_dlist( + self._remote_agents[src_rank], self._nixl_agent.get_xfer_descs(remote_descs, "VRAM"), backends=["UCX"] + ) + xfer = self._nixl_agent.make_prepped_xfer( + "READ", + local_dlist, + list(range(len(local_descs))), + remote_dlist, + list(range(len(remote_descs))), + backends=["UCX"], + ) + selected_backend = self._nixl_agent.query_xfer_backend(xfer) + if selected_backend != "UCX": + raise RuntimeError("NIXL EPLB READ did not select UCX") + self._xfer_cache[cache_key] = (local_dlist, remote_dlist, xfer) + self._used_xfer_cache_keys.add(cache_key) + return self._xfer_cache[cache_key] + except Exception: + self._release_xfers([(local_dlist, remote_dlist, xfer)]) + raise + + def _copy_batch(self, batch) -> None: + remote_entries = defaultdict(list) + push_entries = defaultdict(list) + for layer_index, plan, _, staging in batch: + steps_by_source = defaultdict(list) + for step in plan: + if step.dst_rank == self.global_rank: + steps_by_source[step.src_rank].append(step) + for src_rank, steps in steps_by_source.items(): + entries = [(layer_index, run, staging) for run in self._contiguous_runs(steps)] + if src_rank not in self._same_node_ranks: + remote_entries[src_rank].extend(entries) + # Source rank owns node-local copies. All ranks build the same batch, + # so buffer_index is the receiver's staging depth index on every peer. + for layer_index, plan, buffer_index, _ in batch: + by_destination = defaultdict(list) + for step in plan: + if step.src_rank == self.global_rank and step.dst_rank in self._same_node_ranks: + by_destination[step.dst_rank].append(step) + for dst_rank, steps in by_destination.items(): + push_entries[dst_rank].extend((layer_index, run, buffer_index) for run in self._contiguous_runs(steps)) + for dst_rank, entries in push_entries.items(): + self._push_same_node(dst_rank, entries) + + xfers = [self._get_remote_read(src_rank, entries) for src_rank, entries in remote_entries.items()] + self._wait_xfers(xfers) + self._push_stream.synchronize() + # Before a rank publishes this batch it has completed its outgoing source-pushes and + # incoming UCX READs. The manager's global MIN-ready gate therefore means all transfers + # are complete before any rank commits, without a destination-side GPU wait. + + def _start_transfer_generation(self) -> None: + self._used_xfer_cache_keys.clear() + self._used_push_descriptor_cache_keys.clear() + + def _finish_transfer_generation(self) -> None: + errors = [] + for cache_key in set(self._xfer_cache) - self._used_xfer_cache_keys: + xfer = self._xfer_cache[cache_key] + try: + self._release_xfers([xfer]) + except Exception as exc: + unreleased = getattr(exc, "unreleased_xfers", None) + if unreleased: + self._xfer_cache[cache_key] = unreleased[0] + errors.append(exc) + else: + del self._xfer_cache[cache_key] + for cache_key in set(self._push_descriptor_cache) - self._used_push_descriptor_cache_keys: + del self._push_descriptor_cache[cache_key] + if errors: + raise RuntimeError("NIXL EPLB cache eviction failed") from errors[0] + + def shutdown(self) -> None: + agent = self._nixl_agent + errors = [] + getattr(self, "_used_xfer_cache_keys", set()).clear() + getattr(self, "_used_push_descriptor_cache_keys", set()).clear() + if agent is not None: + for cache_key, xfer in list(self._xfer_cache.items()): + try: + self._release_xfers([xfer]) + except Exception as exc: + unreleased = getattr(exc, "unreleased_xfers", None) + if unreleased: + self._xfer_cache[cache_key] = unreleased[0] + errors.append(exc) + else: + del self._xfer_cache[cache_key] + if errors: + raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] + for remote_name in list(self._remote_agents.values()): + if agent is not None: + try: + agent.remove_remote_agent(remote_name) + except Exception as exc: + errors.append(exc) + self._remote_agents.clear() + self._remote_layouts.clear() + if agent is not None and self._registered_descs is not None: + try: + agent.deregister_memory(self._registered_descs, backends=["UCX"]) + except Exception as exc: + errors.append(exc) + self._registered_descs = None + self._nixl_agent = None + getattr(self, "_ipc_staging", {}).clear() + getattr(self, "_push_descriptor_cache", {}).clear() + if errors: + raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] + + def __del__(self): + try: + self.shutdown() + except Exception: + pass diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index cc7e968193..505588f4ec 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -97,6 +97,36 @@ def get_lightllm_websocket_max_message_size(): @lru_cache(maxsize=None) +def get_prefill_eplb_step_interval(): + """Return the number of prefill forwards between EPLB attempts.""" + interval = int(os.getenv("LIGHTLLM_PREFILL_EPLB_STEP_INTERVAL", 20)) + if interval <= 0: + raise ValueError("LIGHTLLM_PREFILL_EPLB_STEP_INTERVAL must be greater than 0") + return interval + + +@lru_cache(maxsize=None) +def get_eplb_rebalance_gain_threshold() -> float: + """Return the EPLB gain threshold: estimated critical-load reduction ratio; 0.05 means 5%.""" + env_name = "LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD" + raw_value = os.getenv(env_name, "0.05") + value = float(raw_value) + if not 0.0 <= value <= 1.0: + raise ValueError(f"{env_name} must be a ratio between 0.0 and 1.0, got {raw_value!r}") + return value + + +@lru_cache(maxsize=None) +def get_eplb_placement_stickiness() -> float: + """Return the EPLB placement stickiness: a keep-bonus, as a fraction of the mean per-layer expert load.""" + env_name = "LIGHTLLM_EPLB_PLACEMENT_STICKINESS" + raw_value = os.getenv(env_name, "0.1") + value = float(raw_value) + if not 0.0 <= value <= 1.0: + raise ValueError(f"{env_name} must be a ratio between 0.0 and 1.0, got {raw_value!r}") + return value + + def get_triton_autotune_level(): return int(os.getenv("LIGHTLLM_TRITON_AUTOTUNE_LEVEL", 0)) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py new file mode 100644 index 0000000000..38de738a48 --- /dev/null +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -0,0 +1,3360 @@ +import threading +import time +from collections import deque +from contextlib import contextmanager +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_initial_redundant_expert_ids, + build_logical_to_physical_map, + build_logical_to_physical_maps_for_layers, + _estimate_rank_load, + plan_redundant_experts, + select_improving_placements, +) +from lightllm.server.api_cli import make_argument_parser +from lightllm.server.core.objs.start_args_type import StartArgs +from lightllm.server.router.model_infer.infer_batch import g_infer_context +from lightllm.server.router.model_infer.mode_backend import ( + eplb_manager as manager_module, +) +from lightllm.server.router.model_infer.mode_backend import ( + eplb_transfer as transfer_module, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( + deepgemm_impl as deepgemm_module, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + EPLBState, + ExpertParallelState, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( + create_fuse_moe_impl, + FuseMoeMarlin, + FuseMoeTriton, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.base_impl import ( + FuseMoeBaseImpl, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe import ( + fused_moe_weight as fused_weight_module, +) +from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + TransferStep, + align_target_placement, + build_transfer_plan, + commit_staging_rows, + extract_expert_tensors, +) +from lightllm.utils import envs_utils + + +def _test_parallel_state( + *, + eplb=False, + num_logical_experts=128, + world_size=16, + num_redundant_experts_per_rank=1, + route_counter=None, + recording=False, + recorded_sample_count=0, +): + eplb_state = None + if eplb: + if route_counter is None: + route_counter = torch.zeros((2, num_logical_experts), dtype=torch.int64) + initial_layout_world_size = max(world_size, 2) + eplb_state = EPLBState( + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids( + num_logical_experts, + initial_layout_world_size, + num_redundant_experts_per_rank, + ), + logical_to_physical_map=torch.zeros((num_logical_experts, 2), dtype=torch.int32), + logical_replica_count=torch.ones(num_logical_experts, dtype=torch.int32), + route_counter=route_counter, + recording=recording, + recorded_sample_count=recorded_sample_count, + ) + return ExpertParallelState( + num_logical_experts=num_logical_experts, + world_size=world_size, + eplb=eplb_state, + ) + + +def _validated_expert_parallel_state( + *, + eplb=True, + n_routed_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + device="cpu", +): + runtime = None + if eplb: + runtime = EPLBState( + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids( + n_routed_experts, + world_size, + num_redundant_experts_per_rank, + ), + logical_to_physical_map=torch.empty((n_routed_experts, 2), dtype=torch.int32, device=device), + logical_replica_count=torch.ones(n_routed_experts, dtype=torch.int32, device=device), + route_counter=torch.zeros((2, n_routed_experts), dtype=torch.int64, device=device), + ) + return ExpertParallelState( + num_logical_experts=n_routed_experts, + world_size=world_size, + eplb=runtime, + ) + + +def _set_expert_parallel_state(impl, state): + impl.expert_parallel_state = state + impl.eplb = state.eplb + + +def _manual_runtime_rank_load(source_load, placement, node_world_size, alignment): + """Reference the committed runtime logical-to-physical maps on CPU.""" + samples, layers, nodes, num_logical_experts = source_load.shape + ranks, redundant = placement.shape[1:] + num_experts_per_rank = num_logical_experts // ranks + num_physical_experts_per_rank = num_experts_per_rank + redundant + raw = torch.zeros((samples, layers, ranks, num_logical_experts), dtype=torch.float64) + for layer in range(layers): + for source_node in range(nodes): + logical_to_physical, replica_count = build_logical_to_physical_map( + placement[layer], + num_logical_experts, + source_rank=source_node * node_world_size, + node_world_size=node_world_size, + ) + for expert in range(num_logical_experts): + count = int(replica_count[expert].item()) + for physical_id in logical_to_physical[expert, :count].tolist(): + rank = physical_id // num_physical_experts_per_rank + raw[:, layer, rank, expert] += source_load[:, layer, source_node, expert] / count + return (torch.ceil(raw / alignment) * alignment).sum(dim=3) + + +def test_base_call_template_forwards_selection_and_capture_callback(): + class Impl(FuseMoeBaseImpl): + def _select_experts( + self, + input_tensor, + router_logits, + correction_bias, + top_k, + renormalize, + use_grouped_topk, + topk_group, + num_expert_group, + scoring_func, + per_expert_scale=None, + shared_expert_gate=None, + is_prefill=None, + preserve_logical_ids=False, + ): + seen["select"] = {"preserve_logical_ids": preserve_logical_ids} + return "weights", "physical_ids", "logical_ids" + + def _fused_experts( + self, + input_tensor, + w13, + w2, + topk_weights, + topk_ids, + router_logits=None, + is_prefill=None, + ): + seen["fused"] = {"topk_ids": topk_ids} + return "output" + + seen, captured = {}, [] + impl = Impl(4, 0, 1.0, SimpleNamespace()) + result = impl( + "input", + "logits", + "w13", + "w2", + None, + "softmax", + 2, + False, + False, + 0, + 0, + moe_capture_callback=captured.append, + ) + assert result == "output" + assert captured == ["logical_ids"] + assert seen["select"]["preserve_logical_ids"] + assert seen["fused"]["topk_ids"] == "physical_ids" + + +def test_parallel_state_derives_expert_layout(): + state = _validated_expert_parallel_state(eplb=True) + assert state.num_primary_experts_per_rank == 2 + assert state.num_total_physical_experts == 6 + + +def test_factory_selects_all_paths_and_requires_ep_state(): + plain_quant = SimpleNamespace(method_name="none") + marlin_quant = SimpleNamespace(method_name="awq_marlin") + state = _validated_expert_parallel_state(eplb=False) + ep_impl = create_fuse_moe_impl( + n_routed_experts=4, + num_fused_shared_experts=0, + routed_scaling_factor=1.0, + quant_method=plain_quant, + expert_parallel_state=state, + ) + assert isinstance(ep_impl, deepgemm_module.FuseMoeDeepGEMM) + assert ep_impl.expert_parallel_state is state + assert state.eplb is None + assert isinstance( + create_fuse_moe_impl( + n_routed_experts=4, + num_fused_shared_experts=0, + routed_scaling_factor=1.0, + quant_method=plain_quant, + ), + FuseMoeTriton, + ) + assert isinstance( + create_fuse_moe_impl( + n_routed_experts=4, + num_fused_shared_experts=0, + routed_scaling_factor=1.0, + quant_method=marlin_quant, + ), + FuseMoeMarlin, + ) + + +def test_find_fused_moe_weights_discovers_direct_layer_attributes(monkeypatch): + class FakeFusedMoeWeight: + def __init__(self, layer_num, enable_ep_moe=True): + self.layer_num_ = layer_num + self.enable_ep_moe = enable_ep_moe + + monkeypatch.setattr(manager_module, "FusedMoeWeight", FakeFusedMoeWeight) + first = FakeFusedMoeWeight(3) + alternate = FakeFusedMoeWeight(1) + aliased = FakeFusedMoeWeight(2) + disabled = FakeFusedMoeWeight(0, enable_ep_moe=False) + model = SimpleNamespace( + trans_layers_weight=[ + SimpleNamespace(moe_weight=first), + SimpleNamespace(alternate_moe_weight=alternate), + SimpleNamespace(moe_weight=aliased, alternate_moe_weight=aliased), + SimpleNamespace(moe_weight=disabled), + ] + ) + + assert manager_module._find_fused_moe_weights(model) == [alternate, aliased, first] + + +def test_get_eplb_rebalance_gain_threshold_defaults_to_five_percent(monkeypatch): + monkeypatch.delenv("LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD", raising=False) + envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() + + try: + assert envs_utils.get_eplb_rebalance_gain_threshold() == 0.05 + finally: + envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() + + +def test_get_eplb_placement_stickiness_defaults_to_ten_percent(monkeypatch): + monkeypatch.delenv("LIGHTLLM_EPLB_PLACEMENT_STICKINESS", raising=False) + envs_utils.get_eplb_placement_stickiness.cache_clear() + + try: + assert envs_utils.get_eplb_placement_stickiness() == 0.1 + finally: + envs_utils.get_eplb_placement_stickiness.cache_clear() + + +@pytest.mark.parametrize(("configured", "expected"), [("0", 0.0), (".04", 0.04), ("1", 1.0)]) +def test_get_eplb_rebalance_gain_threshold_reads_valid_values(monkeypatch, configured, expected): + monkeypatch.setenv("LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD", configured) + envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() + + try: + assert envs_utils.get_eplb_rebalance_gain_threshold() == expected + finally: + envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() + + +@pytest.mark.parametrize("configured", ["-0.01", "1.01", "nan", "inf"]) +def test_get_eplb_rebalance_gain_threshold_rejects_invalid_values(monkeypatch, configured): + env_name = "LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD" + monkeypatch.setenv(env_name, configured) + envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() + + try: + with pytest.raises(ValueError, match=env_name): + envs_utils.get_eplb_rebalance_gain_threshold() + finally: + envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() + monkeypatch.delenv(env_name, raising=False) + + +def test_eplb_redundant_experts_defaults_per_ep_rank(): + parser = make_argument_parser() + + assert parser.parse_args([]).eplb_num_redundant_experts_per_rank == 2 + assert parser.parse_args(["--eplb_num_redundant_experts_per_rank", "3"]).eplb_num_redundant_experts_per_rank == 3 + assert StartArgs().eplb_num_redundant_experts_per_rank == 2 + + +@pytest.mark.parametrize( + ("num_logical_experts", "num_ranks", "num_redundant_experts_per_rank", "expected"), + [ + (8, 4, 2, [[2, 3], [4, 5], [6, 7], [0, 1]]), + (6, 3, 4, [[2, 3, 4, 5], [4, 5, 0, 1], [0, 1, 2, 3]]), + ], +) +def test_build_initial_redundant_expert_ids( + num_logical_experts, + num_ranks, + num_redundant_experts_per_rank, + expected, +): + actual = build_initial_redundant_expert_ids( + num_logical_experts, + num_ranks, + num_redundant_experts_per_rank, + ) + + assert actual.dtype == torch.int64 + assert actual.shape == (num_ranks, num_redundant_experts_per_rank) + assert torch.equal(actual, torch.tensor(expected, dtype=torch.int64)) + + +def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): + expert_load = torch.tensor( + [ + [100, 90, 80, 70, 60, 50, 40, 30], + [30, 40, 50, 60, 70, 80, 90, 100], + ] + ) + placement = plan_redundant_experts(expert_load, num_ranks=4, num_redundant_experts_per_rank=2) + + for layer_placement in placement: + for rank, expert_ids in enumerate(layer_placement.tolist()): + assert len(expert_ids) == len(set(expert_ids)) + assert all(expert_id // 2 != rank for expert_id in expert_ids) + + +def test_plan_redundant_experts_minimizes_samplewise_aligned_critical_load(): + samples = torch.tensor([[[300, 20, 20, 200]], [[100, 300, 40, 160]]]) + placement = plan_redundant_experts(samples, num_ranks=2, num_redundant_experts_per_rank=1, expert_alignment=128) + candidates = [torch.tensor([[[left], [right]]]) for left in (2, 3) for right in (0, 1)] + + def critical(candidate): + return _estimate_rank_load(samples, candidate, expert_alignment=128).max(dim=2).values.sum() + + assert torch.equal(placement, torch.tensor([[[3], [0]]])) + assert critical(placement) == min(critical(candidate) for candidate in candidates) + + +def test_select_improving_placements_rejects_regressing_layer(): + expert_load = torch.tensor([[8649, 5740, 5002, 3441]]) + current = torch.tensor([[[2], [0]]]) + regressing_candidate = torch.tensor([[[1], [0]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, regressing_candidate, rebalance_gain_threshold=0.05 + ) + + current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() + candidate_ratio = ( + _estimate_rank_load(expert_load, regressing_candidate).max() + / _estimate_rank_load(expert_load, regressing_candidate).mean() + ) + assert current_ratio.item() == pytest.approx(1.1007, abs=1e-4) + assert candidate_ratio.item() == pytest.approx(1.1184, abs=1e-4) + assert not improved.item() + assert torch.equal(selected, current) + + +def test_select_improving_placements_rejects_near_balance_when_gain_is_below_threshold(): + expert_load = torch.tensor([[1, 2, 1, 17]]) + current = torch.tensor([[[3], [0]]]) + candidate = torch.tensor([[[3], [1]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() + candidate_ratio = ( + _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() + ) + assert current_ratio.item() == pytest.approx(1.047619, abs=1e-6) + assert candidate_ratio.item() == pytest.approx(1.0) + assert not improved.item() + assert torch.equal(selected, current) + + +def test_select_improving_placements_accepts_alignment_aware_gain_even_when_current_ranks_are_balanced(): + expert_load = torch.tensor([[100, 129, 100, 129]]) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[3], [1]]]) + + selected, improved, metrics, _before_load, _after_load = select_improving_placements( + expert_load, + current, + candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, + ) + + assert metrics["model_imbalance_ratio"] == pytest.approx(1.0) + assert metrics["candidate_rebalance_gain"] == pytest.approx(0.25) + assert improved.item() + assert torch.equal(selected, candidate) + + +def test_select_improving_placements_rejects_insufficient_rebalance_gain(): + expert_load = torch.tensor([[1, 1, 6, 7]]) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[3], [0]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() + candidate_ratio = ( + _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() + ) + relative_improvement = (current_ratio - candidate_ratio) / current_ratio + assert current_ratio.item() == pytest.approx(1.4) + assert candidate_ratio.item() == pytest.approx(1.333333, abs=1e-6) + assert relative_improvement.item() == pytest.approx(0.047619, abs=1e-6) + assert not improved.item() + assert torch.equal(selected, current) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, + current, + candidate, + rebalance_gain_threshold=0.04, + ) + + assert improved.item() + assert torch.equal(selected, candidate) + + +@pytest.mark.parametrize("rebalance_gain_threshold", [-0.01, 1.01, float("nan"), float("inf")]) +def test_select_improving_placements_rejects_invalid_rebalance_gain_threshold( + rebalance_gain_threshold, +): + expert_load = torch.tensor([[1, 1, 1, 2]]) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[3], [0]]]) + + with pytest.raises(ValueError, match="rebalance_gain_threshold"): + select_improving_placements( + expert_load, + current, + candidate, + rebalance_gain_threshold=rebalance_gain_threshold, + ) + + +def test_select_improving_placements_accepts_sufficient_rebalance_gain(): + expert_load = torch.tensor([[1, 1, 1, 2]]) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[3], [0]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() + candidate_ratio = ( + _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() + ) + relative_improvement = (current_ratio - candidate_ratio) / current_ratio + assert current_ratio.item() == pytest.approx(1.2) + assert candidate_ratio.item() == pytest.approx(1.0) + assert relative_improvement.item() == pytest.approx(1 / 6) + assert improved.item() + assert torch.equal(selected, candidate) + + +def test_select_improving_placements_rejects_raw_improvement_that_does_not_improve_aligned_compute(): + expert_load = torch.tensor([[1, 1, 1, 8]]) + current = torch.tensor([[[2], [0]]]) + raw_improving_candidate = torch.tensor([[[3], [0]]]) + + _, raw_improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, raw_improving_candidate, rebalance_gain_threshold=0.05 + ) + selected, aligned_improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, + current, + raw_improving_candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, + ) + + assert raw_improved.item() + assert not aligned_improved.item() + assert torch.equal(selected, current) + + +def test_estimate_rank_load_aligns_each_sample_before_accumulation(): + samples = torch.tensor([[[20, 0, 0, 0]], [[20, 0, 0, 0]]]) + placement = torch.tensor([[[2], [0]]]) + + per_sample = _estimate_rank_load(samples, placement, expert_alignment=128) + accumulated = _estimate_rank_load(samples.sum(dim=0), placement, expert_alignment=128) + + assert torch.equal(per_sample[:, 0], torch.tensor([[128.0, 128.0], [128.0, 128.0]])) + assert torch.equal(per_sample.sum(dim=0)[0], torch.tensor([256.0, 256.0])) + assert torch.equal(accumulated[0], torch.tensor([128.0, 128.0])) + + +def test_select_improving_placements_rejects_lower_ratio_when_critical_is_unchanged(): + samples = torch.tensor( + [ + [[255, 220, 226, 254]], + [[172, 278, 51, 238]], + [[249, 291, 284, 183]], + ] + ) + current = torch.tensor([[[2], [0]]]) + mean_inflating_candidate = torch.tensor([[[2], [1]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + samples, + current, + mean_inflating_candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, + ) + current_load = _estimate_rank_load(samples, current, expert_alignment=128) + candidate_load = _estimate_rank_load(samples, mean_inflating_candidate, expert_alignment=128) + current_critical = current_load.max(dim=2).values.sum() + candidate_critical = candidate_load.max(dim=2).values.sum() + + assert candidate_load.mean(dim=2).sum() > current_load.mean(dim=2).sum() + assert current_critical == candidate_critical + assert not improved.item() + assert torch.equal(selected, current) + + +def test_select_improving_placements_accepts_five_percent_critical_reduction(): + samples = torch.tensor( + [ + [[13, 352, 348, 141]], + [[287, 175, 236, 179]], + [[316, 99, 266, 353]], + ] + ) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[2], [1]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + samples, current, candidate, rebalance_gain_threshold=0.05, expert_alignment=128 + ) + + assert improved.item() + assert torch.equal(selected, candidate) + + +def test_select_improving_placements_rejects_single_layer_gain_below_model_threshold(): + # Layer 0 becomes better, but layer 1 dominates model critical load. The + # aggregate estimated critical-load reduction gain is below 5%, so neither layer may be changed. + expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 10]]) + current = torch.tensor([[[2], [0]], [[2], [0]]]) + candidate = torch.tensor([[[3], [0]], [[2], [0]]]) + + selected, improved, metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + assert not torch.any(improved) + assert torch.equal(selected, current) + assert metrics["candidate_rebalance_gain"] == pytest.approx(0.5 / 13) + assert metrics["candidate_changed_layer_count"] == 1 + + +def test_select_improving_placements_accepts_only_when_model_gain_reaches_threshold(): + expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 5]]) + current = torch.tensor([[[2], [0]], [[2], [0]]]) + candidate = torch.tensor([[[3], [0]], [[2], [0]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + assert torch.equal(improved, torch.tensor([True, False])) + assert torch.equal(selected, candidate) + + +def test_logical_to_physical_map_has_at_most_one_slot_per_rank(): + redundant_expert_ids = torch.tensor([[2, 3], [0, 1]]) + logical_to_physical, replica_count = build_logical_to_physical_map(redundant_expert_ids, num_logical_experts=4) + + assert logical_to_physical.shape == (4, 2) + assert torch.equal(replica_count, torch.full((4,), 2, dtype=torch.int64)) + + +def test_logical_to_physical_map_requires_expert_count_divisible_by_rank_count_without_source_rank(): + redundant_expert_ids = torch.tensor([[0], [1]]) + + with pytest.raises(AssertionError): + build_logical_to_physical_map(redundant_expert_ids, num_logical_experts=5) + + +def test_logical_to_physical_map_prefers_source_node_replicas(): + # Four ranks, two ranks per node, two primary experts/rank and one + # redundant slot/rank. Expert 0 is primary on rank 0 and replicated on + # rank 2 (the other node). + redundant = torch.tensor([[4], [5], [0], [1]], dtype=torch.int64) + rank0_map, rank0_count = build_logical_to_physical_map( + redundant, num_logical_experts=8, source_rank=0, node_world_size=2 + ) + rank1_map, rank1_count = build_logical_to_physical_map( + redundant, num_logical_experts=8, source_rank=1, node_world_size=2 + ) + fallback_redundant = torch.tensor([[4], [5], [0], [1], [2], [3]], dtype=torch.int64) + rank4_map, rank4_count = build_logical_to_physical_map(fallback_redundant, 12, source_rank=4, node_world_size=2) + + assert rank0_count[0].item() == rank1_count[0].item() == 1 + assert torch.equal(rank0_map[0, :1], torch.tensor([0])) + assert torch.equal(rank1_map[0, :1], torch.tensor([0])) + assert rank4_count[0].item() == 2 + assert set(rank4_map[0, :2].tolist()) == {0, 8} + + +def test_source_node_local_maps_fall_back_to_global_replicas(): + redundant = torch.tensor([[4], [5], [0], [1]], dtype=torch.int64) + maps = [build_logical_to_physical_map(redundant, 8, source_rank=rank, node_world_size=2) for rank in range(4)] + assert maps[0][1][0].item() == maps[1][1][0].item() == 1 + assert maps[2][1][0].item() == maps[3][1][0].item() == 1 + assert maps[0][0][0, 0].item() == maps[1][0][0, 0].item() == 0 + assert maps[2][0][0, 0].item() == maps[3][0][0, 0].item() == 8 + + +def test_source_rank_rotates_selected_replica_order_without_changing_copies(): + # Expert 0 is primary on rank 0 and redundant on rank 1, so both ranks + # on node 0 have the same two local copies. Their source-rank phases + # must differ while their selected set/count remain identical. + redundant = torch.tensor([[1], [0], [3], [2]], dtype=torch.int64) + rank0_map, rank0_count = build_logical_to_physical_map(redundant, 4, source_rank=0, node_world_size=2) + rank1_map, rank1_count = build_logical_to_physical_map(redundant, 4, source_rank=1, node_world_size=2) + + assert rank0_count[0].item() == rank1_count[0].item() == 2 + assert set(rank0_map[0, :2].tolist()) == set(rank1_map[0, :2].tolist()) == {0, 3} + assert torch.equal(rank1_map[0, :2], torch.tensor([3, 0], dtype=torch.int32)) + + +@pytest.mark.parametrize( + "source_rank,node_world_size", + [(None, None), (0, 2), (1, 2), (2, 2), (3, 2)], +) +def test_logical_to_physical_maps_for_layers_match_single_layer_api(source_rank, node_world_size): + placements_by_layer = torch.tensor( + [ + [[4, 5], [0, 1], [0, 1], [2, 3]], + [[6, 7], [0, 1], [0, 1], [2, 3]], + [[4, 5], [0, 1], [0, 1], [2, 3]], + ], + dtype=torch.int64, + ) + + maps_by_layer, counts_by_layer = build_logical_to_physical_maps_for_layers( + placements_by_layer, + num_logical_experts=8, + source_rank=source_rank, + node_world_size=node_world_size, + ) + expected_by_layer = [ + build_logical_to_physical_map( + placement, + num_logical_experts=8, + source_rank=source_rank, + node_world_size=node_world_size, + ) + for placement in placements_by_layer + ] + + assert maps_by_layer.shape == (3, 8, 4) + assert counts_by_layer.shape == (3, 8) + assert maps_by_layer.dtype == counts_by_layer.dtype == torch.int32 + assert torch.equal(maps_by_layer, torch.stack([item[0] for item in expected_by_layer])) + assert torch.equal(counts_by_layer, torch.stack([item[1] for item in expected_by_layer])) + if source_rank is None: + # expert 0 的主副本在 rank 0,且 rank 1、2 都有一个冗余副本;第 1、2 + # 个冗余副本必须分别写入映射表的第 1、2 列,而不能互相覆盖。 + assert torch.equal(counts_by_layer[:, 0], torch.tensor([3, 3, 3], dtype=torch.int32)) + assert torch.equal( + maps_by_layer[:, 0, :3], + torch.tensor([[0, 6, 10], [0, 6, 10], [0, 6, 10]], dtype=torch.int32), + ) + positions = torch.arange(maps_by_layer.shape[-1]).view(1, 1, -1) + valid = positions < counts_by_layer.unsqueeze(-1) + assert torch.all(maps_by_layer[valid] >= 0) + assert torch.all(maps_by_layer[~valid] == -1) + + +def test_plan_redundant_experts_prefers_first_replica_on_new_node(): + # One redundant slot per rank leaves legal alternatives on both nodes; + # topology preference therefore puts every first replica away from its + # primary node before considering same-node duplicates. + placement = plan_redundant_experts( + torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]), + num_ranks=4, + num_redundant_experts_per_rank=1, + node_world_size=2, + ) + for rank, expert in enumerate(placement[0, :, 0].tolist()): + assert expert // 2 // 2 != rank // 2 + + +def test_plan_redundant_experts_single_node_matches_default_behavior(): + load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]) + default = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1) + single_node = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1, node_world_size=4) + assert torch.equal(single_node, default) + + +def test_source_node_estimate_matches_local_first_runtime_replica_sharing(): + # Expert 0 has a copy on each node. Node 0 and node 1 issue unequal + # traffic, so collapsing them before planning produces the wrong result. + placement = torch.tensor([[[4], [5], [0], [1]]], dtype=torch.int64) + source_load = torch.zeros((1, 1, 2, 8), dtype=torch.int64) + source_load[0, 0, 0, 0] = 256 + source_load[0, 0, 1, 0] = 128 + + predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) + runtime = _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128) + collapsed_global = _estimate_rank_load(source_load.sum(dim=2), placement, expert_alignment=128) + + assert torch.equal(predicted, runtime) + assert torch.equal(predicted[0, 0], torch.tensor([256.0, 0.0, 128.0, 0.0])) + assert not torch.equal(predicted, collapsed_global) + + +def test_source_node_planner_constraints_and_real_critical_improvement(): + source_load = torch.tensor( + [ + [ + [ + [697, 451, 383, 536, 349, 404, 854, 425], + [861, 103, 166, 612, 444, 263, 910, 392], + ] + ], + [ + [ + [944, 457, 338, 108, 63, 525, 48, 216], + [439, 117, 837, 550, 833, 201, 729, 5], + ] + ], + [ + [ + [749, 159, 18, 723, 12, 700, 419, 51], + [112, 135, 8, 840, 40, 970, 90, 683], + ] + ], + ], + dtype=torch.int64, + ) + initial = build_initial_redundant_expert_ids(8, 4, 1).unsqueeze(0) + planned = plan_redundant_experts(source_load, 4, 1, expert_alignment=128, node_world_size=2) + + for rank, experts in enumerate(planned[0].tolist()): + assert len(experts) == len(set(experts)) == 1 + assert experts[0] // 2 != rank + + before = _estimate_rank_load(source_load, initial, expert_alignment=128, node_world_size=2) + after = _estimate_rank_load(source_load, planned, expert_alignment=128, node_world_size=2) + manual_before = _manual_runtime_rank_load(source_load, initial, 2, 128) + manual_after = _manual_runtime_rank_load(source_load, planned, 2, 128) + assert torch.equal(before, manual_before) + assert torch.equal(after, manual_after) + assert after.max(dim=2).values.sum() < before.max(dim=2).values.sum() + + +def test_source_node_select_uses_the_same_runtime_critical_prediction(): + source_load = torch.zeros((2, 1, 2, 8), dtype=torch.int64) + source_load[:, 0, 0, 0] = torch.tensor([1024, 768]) + source_load[:, 0, 1, 6] = torch.tensor([896, 1024]) + current = torch.tensor([[[2], [4], [6], [0]]], dtype=torch.int64) + candidate = plan_redundant_experts(source_load, 4, 1, expert_alignment=128, node_world_size=2) + selected, _improved, _metrics, _before_load, _after_load = select_improving_placements( + source_load, + current, + candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, + node_world_size=2, + ) + assert torch.equal( + _estimate_rank_load(source_load, selected, 128, 2), + _manual_runtime_rank_load(source_load, selected, 2, 128), + ) + + +def _count_moved_slots(current: torch.Tensor, target: torch.Tensor) -> int: + """Count rank rows gaining an expert: one migrated row per new expert id.""" + moved = 0 + for layer in range(current.shape[0]): + for rank in range(current.shape[1]): + moved += len(set(target[layer, rank].tolist()) - set(current[layer, rank].tolist())) + return moved + + +def test_sticky_plan_reproduces_current_when_load_unchanged(): + generator = torch.Generator().manual_seed(7) + load = torch.randint(1, 1000, (3, 16, 32), generator=generator) + placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) + + replanned = plan_redundant_experts( + load, + num_ranks=4, + num_redundant_experts_per_rank=2, + current_placement=placement, + stickiness=0.1, + ) + + assert torch.equal(replanned, placement) + for layer in range(placement.shape[0]): + assert build_transfer_plan(placement[layer], replanned[layer], 32, 4, 4) == [] + + +def test_sticky_plan_bounded_moves_under_small_perturbation(): + generator = torch.Generator().manual_seed(11) + load = torch.randint(100, 1000, (4, 16, 32), generator=generator) + placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) + noise = torch.rand((4, 16, 32), generator=generator) * 0.1 + 0.95 + perturbed = (load.double() * noise).round().to(torch.int64) + + sticky = plan_redundant_experts(perturbed, 4, 2, current_placement=placement, stickiness=0.1) + free = plan_redundant_experts(perturbed, 4, 2) + + sticky_moves = _count_moved_slots(placement, sticky) + free_moves = _count_moved_slots(placement, free) + assert sticky_moves <= placement.numel() // 4 + assert sticky_moves < free_moves + + def critical(candidate): + return _estimate_rank_load(perturbed, candidate).max(dim=2).values.sum() + + assert critical(sticky) <= critical(free) * 1.1 + + +def test_sticky_plan_still_churns_under_phase_shift(): + layers, experts = 8, 32 + before = torch.full((layers, experts), 10, dtype=torch.int64) + after = torch.full((layers, experts), 10, dtype=torch.int64) + offsets = torch.arange(4) + for layer in range(layers): + before[layer, (4 * layer + offsets) % experts] = 5000 + after[layer, (4 * layer + 16 + offsets) % experts] = 5000 + placement = plan_redundant_experts(before, num_ranks=4, num_redundant_experts_per_rank=2) + + replanned = plan_redundant_experts( + after, + num_ranks=4, + num_redundant_experts_per_rank=2, + current_placement=placement, + stickiness=0.1, + ) + + assert _count_moved_slots(placement, replanned) > placement.numel() // 2 + + +def test_transfer_plan_slot_permutation_is_free(): + current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) + target = torch.tensor([[5, 4], [7, 6], [1, 0], [3, 2]]) + + assert torch.equal(align_target_placement(current, target), current) + assert build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) == [] + + +def test_align_target_placement_keeps_retained_experts_in_live_slots(): + current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) + target = torch.tensor([[5, 6], [7, 0], [1, 0], [3, 2]]) + + canonical = align_target_placement(current, target) + + assert torch.equal(canonical, torch.tensor([[6, 5], [0, 7], [0, 1], [2, 3]])) + + +def test_canonical_placement_keeps_transfer_rows_and_published_map_consistent(): + num_logical_experts = 8 + world_size = 4 + num_redundant_slots_per_rank = 2 + num_experts_per_rank = num_logical_experts // world_size + num_physical_experts_per_rank = num_experts_per_rank + num_redundant_slots_per_rank + current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) + target = torch.tensor([[5, 6], [7, 0], [1, 0], [3, 2]]) + canonical = align_target_placement(current, target) + plan = build_transfer_plan(current, canonical, num_logical_experts, world_size, node_world_size=2) + + # Label every current physical row by its resident logical expert, then + # apply the transfer plan from a frozen source snapshot just as staging + # copies do before the destination rows are published. + source_rows = [ + list(range(rank * num_experts_per_rank, (rank + 1) * num_experts_per_rank)) + current[rank].tolist() + for rank in range(world_size) + ] + live_rows = [row.copy() for row in source_rows] + for step in plan: + live_rows[step.dst_rank][num_experts_per_rank + step.dst_slot] = source_rows[step.src_rank][step.src_local_row] + + logical_to_physical, replica_count = build_logical_to_physical_map(canonical, num_logical_experts) + for logical_expert, count in enumerate(replica_count.tolist()): + for physical_id in logical_to_physical[logical_expert, :count].tolist(): + rank, row = divmod(physical_id, num_physical_experts_per_rank) + assert live_rows[rank][row] == logical_expert + + +def test_plan_and_broadcast_publishes_canonical_placement(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.global_rank = 0 + manager.world_size = 4 + manager.node_world_size = 2 + manager.num_logical_experts = 8 + manager.num_redundant_experts_per_rank = 2 + manager.current_placement = torch.tensor([[[4, 5], [6, 7], [0, 1], [2, 3]]]) + manager.placement_stickiness = 0.1 + manager.rebalance_gain_threshold = 0.05 + manager.evaluation_group = object() + candidate = torch.tensor([[[5, 4], [7, 6], [1, 0], [3, 2]]]) + broadcasts = [] + + def fixed_selector(*_args, **_kwargs): + rank_load = torch.full((1, 4), 100.0) + return candidate.clone(), torch.tensor([True]), {}, rank_load, rank_load + + def record_broadcast(result_list, **_kwargs): + broadcasts.append(result_list[0]) + + monkeypatch.setattr(manager_module, "plan_redundant_experts", lambda *_args, **_kwargs: candidate.clone()) + monkeypatch.setattr(manager_module, "select_improving_placements", fixed_selector) + monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) + + result = manager._plan_and_broadcast(torch.full((1, 1, 2, 8), 100, dtype=torch.int64)) + + assert torch.equal(result["placement"], manager.current_placement) + assert broadcasts and torch.equal(broadcasts[0]["placement"], manager.current_placement) + + +def test_stickiness_zero_matches_legacy(): + generator = torch.Generator().manual_seed(17) + load = torch.randint(1, 1000, (2, 8, 16), generator=generator) + legacy = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) + unrelated = build_initial_redundant_expert_ids(16, 4, 2).unsqueeze(0).expand(8, -1, -1).clone() + + replanned = plan_redundant_experts( + load, + num_ranks=4, + num_redundant_experts_per_rank=2, + current_placement=unrelated, + stickiness=0.0, + ) + + assert torch.equal(replanned, legacy) + + +def test_plan_and_broadcast_propagates_rank_zero_error_after_existing_broadcast(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.global_rank = 0 + manager.world_size = 2 + manager.node_world_size = 2 + manager.num_logical_experts = 8 + manager.num_redundant_experts_per_rank = 2 + manager.current_placement = torch.tensor([[[4, 5], [6, 7], [0, 1], [2, 3]]]) + manager.placement_stickiness = 0.1 + manager.rebalance_gain_threshold = 0.05 + manager.evaluation_group = object() + broadcasted = [] + + def broken_planner(*_args, **_kwargs): + raise RuntimeError("planner boom") + + def record_broadcast(result_list, **_kwargs): + broadcasted.append(result_list[0]) + + monkeypatch.setattr(manager_module, "plan_redundant_experts", broken_planner) + monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) + + with pytest.raises(RuntimeError, match="EPLB planner failed on rank zero") as exc_info: + manager._plan_and_broadcast(torch.full((1, 1, 1, 8), 100, dtype=torch.int64)) + + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert str(exc_info.value.__cause__) == "planner boom" + assert broadcasted == [{"kind": "error", "message": "RuntimeError: planner boom"}] + + +def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_evaluates_step_twenty( + monkeypatch, +): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + counter = torch.tensor([[10, 20, 30, 40], [0, 0, 0, 0]], dtype=torch.int64) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=1, + num_logical_experts=4, + world_size=1, + ), + }, + )() + ] + manager.in_flight = False + manager.prefill_steps = 15 + manager.step_interval = 20 + manager.sampling_interval = manager.step_interval + manager.evaluation_group = object() + manager.num_logical_experts = 4 + manager.global_rank = 1 + manager.evaluation_in_flight = False + manager._sampling_pending = False + manager._steady_collection_end_step = None + manager._continuous_collection_start_step = None + manager._continuous_collection_end_step = None + + recordings, resets, started = [], [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) + monkeypatch.setattr( + manager_module.torch.cuda, + "synchronize", + lambda: pytest.fail("step must not synchronize CUDA"), + ) + + manager.step() + assert manager.prefill_steps == 16 + assert recordings == [True] + assert resets == [True] + assert manager._sampling_pending + assert manager._steady_collection_end_step == 20 + + for _ in range(3): + manager.step() + assert manager.prefill_steps == 19 + assert recordings == [True] + assert started == [] + + manager.step() + assert manager.prefill_steps == 20 + assert not manager._sampling_pending + assert started == [True] + + +def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.evaluation_in_flight = False + manager.prefill_steps = 0 + manager.step_interval = 20 + manager.sampling_interval = 3 + manager._sampling_pending = False + manager._steady_collection_end_step = None + manager._continuous_collection_start_step = None + manager._continuous_collection_end_step = None + recordings, resets, started = [], [], [] + manager._set_recording = lambda enabled: recordings.append((manager.prefill_steps, enabled)) + manager._reset_recorded_samples = lambda: resets.append(manager.prefill_steps) + monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(manager.prefill_steps)) + + manager._prepare_next_sampling_window() + assert resets == [0] + assert recordings == [(0, True)] + assert manager._steady_collection_end_step == 3 + assert manager._sampling_pending + + manager.step() + manager.step() + assert started == [] + manager.step() + assert started == [3] + + +def test_eplb_step_does_not_start_a_second_evaluation_while_one_is_pending(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + counter = torch.tensor([[10, 20, 30, 40], [0, 0, 0, 0]], dtype=torch.int64) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=1, + num_logical_experts=4, + world_size=1, + ), + }, + )() + ] + manager.in_flight = False + manager.prefill_steps = 1 + manager.step_interval = 2 + manager.sampling_interval = manager.step_interval + manager.evaluation_group = object() + manager.num_logical_experts = 4 + manager.global_rank = 1 + manager.evaluation_in_flight = True + + started = [] + monkeypatch.setattr(manager, "_poll_evaluation", lambda: True) + monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) + + manager.step() + + assert manager.prefill_steps == 1 + assert started == [] + + +def test_evaluation_no_improvement_logs_model_fields_without_reopening_interval_window( + monkeypatch, +): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 2, + } + manager.global_rank = 0 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.weights = [] + manager._eplb_states = [] + recordings, logs = [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) + + assert not manager._poll_evaluation() + assert recordings == [False] + assert "model_imbalance_ratio" in logs[0][0] + assert "candidate_rebalance_gain" in logs[0][0] + assert "candidate_changed_layer_count" in logs[0][0] + assert "actual_changed_layer_count" in logs[0][0] + assert "next_sampling_interval" in logs[0][0] + assert manager.sampling_interval == 80 + + +def test_interval_one_rearms_after_evaluation_but_never_evaluates_empty_counter( + monkeypatch, +): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.prefill_steps = 1 + manager.step_interval = 1 + manager.sampling_interval = 1 + manager._sampling_pending = False + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 2, + } + manager.global_rank = 1 + manager.weights = [] + manager._eplb_states = [] + recordings, starts = [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._start_evaluation = lambda: starts.append(True) + manager._evaluation_ready_on_all_ranks = lambda: True + + manager.poll() + + # The no-improvement backoff changes interval 1 to 4. The clamped + # steady window arms immediately but still waits for boundary step 5. + assert recordings == [True] + assert manager._sampling_pending + assert manager._steady_collection_end_step == 5 + assert starts == [] + assert manager.prefill_steps == 1 + + for _ in range(3): + manager.step() + assert starts == [] + manager.step() + assert starts == [True] + + +def test_evaluation_worker_error_is_raised_by_main_thread(): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = RuntimeError("planner failed") + manager._evaluation_result = None + manager._evaluation_thread = DoneThread() + + with pytest.raises(RuntimeError, match="planner failed"): + manager._poll_evaluation() + + +def test_evaluation_state_is_cleared_before_second_round(monkeypatch): + class DoneThread: + def join(self): + pass + + class PendingThread: + def __init__(self, **_kwargs): + self.started = False + + def start(self): + self.started = True + + def join(self): + pytest.fail("pending worker must not be joined") + + class Event: + def record(self, _stream): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 1, + } + manager.global_rank = 1 + manager.prefill_steps = 0 + manager.step_interval = 1 + manager.sampling_interval = 1 + manager.weights = [] + manager._eplb_states = [] + manager._set_recording = lambda _enabled: None + monkeypatch.setattr(manager_module.threading, "Thread", PendingThread) + monkeypatch.setattr(manager_module.torch.cuda, "Event", Event) + monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: object()) + + assert not manager._poll_evaluation() + assert manager._evaluation_result is None + assert manager._evaluation_error is None + + manager._start_evaluation() + assert manager.evaluation_in_flight + assert manager._evaluation_result is None + assert manager._poll_evaluation() # New worker has not produced a result. + + +def test_manager_collects_recent_ring_samples_in_chronological_order(monkeypatch): + counters = [ + torch.tensor([[10, 11], [20, 21], [30, 31]], dtype=torch.int64), + torch.tensor([[40, 41], [50, 51], [60, 61]], dtype=torch.int64), + ] + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=5, + num_logical_experts=2, + world_size=1, + ), + }, + )() + for counter in counters + ] + manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] + manager.evaluation_group = object() + metadata_sizes = [] + + def all_reduce(metadata, **_kwargs): + metadata_sizes.append(metadata.numel()) + + monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) + + samples = manager._collect_local_samples() + + assert metadata_sizes == [4] + assert torch.equal( + samples, + torch.tensor( + [ + [[30, 31], [60, 61]], + [[10, 11], [40, 41]], + [[20, 21], [50, 51]], + ], + dtype=torch.int64, + ), + ) + + +def test_manager_collects_only_two_recent_sparse_samples(monkeypatch): + counters = [ + torch.tensor([[10, 11], [20, 21], [30, 31], [40, 41]], dtype=torch.int64), + torch.tensor([[50, 51], [60, 61], [70, 71], [80, 81]], dtype=torch.int64), + ] + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=2, + num_logical_experts=2, + world_size=1, + ), + }, + )() + for counter in counters + ] + manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] + manager.evaluation_group = object() + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **kwargs: None) + + samples = manager._collect_local_samples() + + assert torch.equal( + samples, + torch.tensor([[[10, 11], [50, 51]], [[20, 21], [60, 61]]], dtype=torch.int64), + ) + + +def test_eplb_counter_capacity_covers_default_dense_interval(monkeypatch): + args = type( + "Args", + (), + {"enable_prefill_eplb": True, "eplb_num_redundant_experts_per_rank": 2}, + )() + weight = object.__new__(fused_weight_module.FusedMoeWeight) + weight.n_routed_experts = 4 + weight.global_world_size = 2 + weight.global_rank_ = 0 + weight.enable_ep_moe = True + monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) + monkeypatch.setattr(fused_weight_module, "get_prefill_eplb_step_interval", lambda: 20) + monkeypatch.setattr(fused_weight_module, "get_node_world_size", lambda: 2) + monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) + original_zeros = torch.zeros + + def cpu_zeros(*shape, **kwargs): + kwargs.pop("device", None) + return original_zeros(*shape, **kwargs) + + monkeypatch.setattr(fused_weight_module.torch, "zeros", cpu_zeros) + + weight._init_expert_parallel_state() + + assert weight.expert_parallel_state.eplb.route_counter.shape == (40, 4) + + +def test_steady_sampling_resets_fixed_ring_without_retained_history(): + counter = torch.ones((8, 4), dtype=torch.int64) + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + state = _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=99, + num_logical_experts=4, + world_size=1, + ).eplb + manager._eplb_states = [state] + + manager._reset_recorded_samples() + manager._reset_recorded_samples() + + assert state.route_counter.shape == (8, 4) + assert torch.count_nonzero(state.route_counter) == 0 + assert state.recorded_sample_count == 0 + assert not hasattr(manager, "_retained_local_samples") + assert not hasattr(manager, "_sample_history") + + +def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): + args = type( + "Args", + (), + {"enable_prefill_eplb": False, "eplb_num_redundant_experts_per_rank": 2}, + )() + weight = object.__new__(fused_weight_module.FusedMoeWeight) + weight.n_routed_experts = 4 + weight.global_world_size = 2 + weight.global_rank_ = 0 + weight.enable_ep_moe = True + monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) + + weight._init_expert_parallel_state() + + assert weight.expert_parallel_state is not None + assert weight.expert_parallel_state.eplb is None + assert weight.expert_parallel_state.num_total_physical_experts == weight.n_routed_experts + assert weight._initial_redundant_expert_ids == [] + assert not hasattr(weight, "route_counter") + assert not hasattr(weight, "routed_expert_counter_tensor") + + +def test_manager_evaluation_collective_preserves_source_node_axis(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=torch.zeros((1, 4), dtype=torch.int64), + num_logical_experts=4, + world_size=1, + ) + }, + )() + ] + manager._eplb_states = [manager.weights[0].expert_parallel_state.eplb] + manager.global_rank = 2 + manager.world_size = 4 + manager.node_world_size = 2 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.num_logical_experts = 4 + manager.num_redundant_experts_per_rank = 1 + manager.current_placement = build_initial_redundant_expert_ids(4, 4, 1).unsqueeze(0) + manager.evaluation_group = object() + manager._evaluation_lock = threading.Lock() + manager._evaluation_result = None + manager._evaluation_error = None + manager._continuous_collection_end_step = None + local = torch.full((1, 1, 4), 100, dtype=torch.int64) + manager._collect_local_samples = lambda: local + seen = {} + + def all_reduce(tensor, **kwargs): + seen["before"] = tensor.clone() + seen["group"] = kwargs["group"] + # Simulate source node 0's contribution from the other ranks. + tensor[:, :, 0].fill_(100) + + monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) + monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) + + def plan_and_broadcast(global_load): + seen["global_load"] = global_load.clone() + return {"kind": "insufficient"} + + manager._plan_and_broadcast = plan_and_broadcast + + manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) + + assert seen["group"] is manager.evaluation_group + expected_local = local + assert seen["before"].shape == (1, 1, 2, 4) + assert torch.equal(seen["before"][:, :, 0], torch.zeros_like(expected_local)) + assert torch.equal(seen["before"][:, :, 1], expected_local) + assert torch.equal(seen["global_load"][:, :, 0], torch.full_like(expected_local, 100)) + assert torch.equal(seen["global_load"][:, :, 1], expected_local) + assert manager._evaluation_error is None + assert manager._evaluation_result["recorded_sample_count"] == 1 + assert manager._evaluation_result["sample_window_steps"] == 4 + + +def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_call(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=torch.zeros((1, 4), dtype=torch.int64), + num_logical_experts=4, + world_size=4, + ) + }, + )() + for _ in range(3) + ] + manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] + manager.global_rank = 1 + manager.world_size = 4 + manager.node_world_size = 2 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.num_logical_experts = 4 + manager.num_redundant_experts_per_rank = 1 + manager.current_placement = build_initial_redundant_expert_ids(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() + manager.evaluation_group = object() + manager._evaluation_lock = threading.Lock() + manager._evaluation_result = None + manager._evaluation_error = None + manager._continuous_collection_end_step = None + manager._collect_local_samples = lambda: torch.full((1, 3, 4), 100, dtype=torch.int64) + planned_placement = torch.tensor( + [ + [[3], [0], [1], [2]], + [[2], [3], [0], [1]], + [[1], [2], [3], [0]], + ], + dtype=torch.int64, + ) + manager._plan_and_broadcast = lambda _global_load: { + "kind": "planned", + "placement": planned_placement, + "improved": torch.tensor([True, False, True]), + } + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) + monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) + calls = [] + original_build_maps_for_layers = manager_module.build_logical_to_physical_maps_for_layers + + def build_maps_for_layers(*args, **kwargs): + calls.append(args[0].shape) + return original_build_maps_for_layers(*args, **kwargs) + + monkeypatch.setattr( + manager_module, + "build_logical_to_physical_maps_for_layers", + build_maps_for_layers, + ) + + manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) + + assert manager._evaluation_error is None + assert calls == [torch.Size([2, 4, 1])] + metadata = manager._evaluation_result["metadata"] + assert metadata[1] is None + assert [layer_index for layer_index, _plan in manager._evaluation_result["layer_plans"]] == [0, 2] + for layer_index in (0, 2): + item = metadata[layer_index] + expected = build_logical_to_physical_map( + planned_placement[layer_index], + 4, + source_rank=manager.global_rank, + node_world_size=manager.node_world_size, + ) + assert torch.equal(item[0], expected[0]) + assert torch.equal(item[1], expected[1]) + + +def test_decode_dispatch_keeps_logical_ids_and_uses_logical_expert_count(monkeypatch): + class Buffer: + def low_latency_dispatch(self, **kwargs): + calls.append(kwargs) + return "recv", "masked", "handle", "event", "hook" + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.quant_method = type("Quant", (), {"method_name": "fp8"})() + impl.n_routed_experts = 128 + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + impl._select_experts = lambda **_kwargs: ( + torch.ones((1, 2)), + torch.tensor([[0, 127]], dtype=torch.int32), + torch.tensor([[0, 127]], dtype=torch.int32), + ) + calls = [] + monkeypatch.setattr( + deepgemm_module, + "get_deepep_num_max_dispatch_tokens_per_rank_decode", + lambda: 16, + ) + monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_low_latency_buffer", Buffer()) + + result = impl.low_latency_dispatch( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + False, + 2, + False, + 0, + 0, + "softmax", + ) + + assert result[2].tolist() == [[0, 127]] + assert calls[0]["num_experts"] == 128 + + +def test_decode_select_does_not_clone_or_map_eplb_topk_ids(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import topk_select + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.routed_scaling_factor = 1.0 + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + topk_ids = torch.tensor([[3, 127]], dtype=torch.int32) + monkeypatch.setattr(topk_select, "select_experts", lambda **_kwargs: (torch.ones((1, 2)), topk_ids)) + _, selected, origin = impl._select_experts( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + 2, + False, + False, + 0, + 0, + "softmax", + is_prefill=False, + ) + + assert selected.data_ptr() == origin.data_ptr() + + +def test_eplb_prefill_uses_single_fused_path_for_global_topk(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.routed_scaling_factor = 1.0 + impl.quant_method = object() + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + physical_ids = torch.tensor([[130, 131]], dtype=torch.long) + calls = [] + + def fused_topk(**kwargs): + calls.append(kwargs) + return torch.ones((1, 2)), physical_ids, None + + monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) + monkeypatch.setattr(deepgemm_module, "quantize_fused_experts_input", lambda *_args: "qinput") + + weights, topk_idx, qinput = impl.select_experts_and_quant_input( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + object(), + False, + 2, + False, + 0, + 0, + "softmax", + ) + + assert weights.tolist() == [[1.0, 1.0]] + assert topk_idx is physical_ids + assert topk_idx.dtype is torch.long + assert qinput == "qinput" + assert not calls[0]["use_grouped_topk"] + assert not calls[0]["return_logical_ids"] + + +def test_eplb_prefill_dispatch_consumes_physical_ids_and_event(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk + + class Buffer: + def dispatch(self, _qinput, **kwargs): + calls.append(kwargs) + return ( + (torch.empty((4, 2)),), + "recv_idx", + "recv_weight", + SimpleNamespace(num_recv_tokens_per_expert_list=[4]), + SimpleNamespace(current_stream_wait=lambda: None), + ) + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.routed_scaling_factor = 1.0 + impl.quant_method = object() + state = _test_parallel_state( + eplb=True, + route_counter=torch.zeros((3, 128), dtype=torch.int64), + recording=True, + ) + _set_expert_parallel_state(impl, state) + impl.ep_balance_counters = None + calls, fused_calls = [], [] + physical_ids = torch.tensor([[130, 131]], dtype=torch.long) + + def fused_topk(**kwargs): + fused_calls.append(kwargs) + return torch.ones((1, 2)), physical_ids, None + + monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) + monkeypatch.setattr(deepgemm_module, "quantize_fused_experts_input", lambda *_args: "qinput") + monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_buffer", Buffer()) + monkeypatch.setattr( + deepgemm_module, + "get_deepep_num_max_dispatch_tokens_per_rank_prefill", + lambda: 16, + ) + monkeypatch.setattr(deepgemm_module, "get_ep_num_sms", lambda: 8) + + weights, topk_idx, qinput = impl.select_experts_and_quant_input( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + object(), + True, + 2, + False, + 1, + 8, + "sigmoid", + ) + caller_event = object() + impl.dispatch( + qinput, + topk_idx, + weights, + overlap_event=caller_event, + ) + + assert topk_idx is physical_ids + assert len(fused_calls) == 1 + assert fused_calls[0]["sample_index"] == 0 + assert fused_calls[0]["record_load"] + assert state.eplb.recorded_sample_count == 1 + assert calls[0]["topk_idx"] is physical_ids + assert calls[0]["topk_idx"].dtype is torch.long + assert calls[0]["previous_event"] is caller_event + + +def test_prefill_dispatch_preserves_event(monkeypatch): + class Buffer: + def dispatch(self, _qinput, **kwargs): + calls.append(kwargs) + return ( + (torch.empty((4, 2)),), + "recv_idx", + "recv_weight", + SimpleNamespace(num_recv_tokens_per_expert_list=[4]), + SimpleNamespace(current_stream_wait=lambda: None), + ) + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + impl.ep_balance_counters = None + calls = [] + caller_event = object() + monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_buffer", Buffer()) + monkeypatch.setattr( + deepgemm_module, + "get_deepep_num_max_dispatch_tokens_per_rank_prefill", + lambda: 16, + ) + monkeypatch.setattr(deepgemm_module, "get_ep_num_sms", lambda: 8) + + impl.dispatch( + "qinput", + torch.tensor([[1, 2]], dtype=torch.long), + torch.ones((1, 2)), + caller_event, + ) + + assert calls[0]["previous_event"] is caller_event + assert calls[0]["topk_idx"].dtype is torch.long + + +def test_deepgemm_constructor_configures_eplb(): + state = _validated_expert_parallel_state() + impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace(), expert_parallel_state=state) + assert impl.expert_parallel_state is state + + +def test_prefill_eplb_returns_requested_logical_ids(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.routed_scaling_factor = 1.0 + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + + def fused_topk(**kwargs): + assert kwargs["return_logical_ids"] + return ( + torch.ones((1, 2)), + torch.tensor([[13, 14]], dtype=torch.int32), + torch.tensor([[3, 4]], dtype=torch.int32), + ) + + monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) + _, physical_ids, logical_ids = impl._select_experts( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + 2, + False, + False, + 0, + 0, + "softmax", + is_prefill=True, + preserve_logical_ids=True, + ) + + assert physical_ids.tolist() == [[13, 14]] + assert logical_ids.tolist() == [[3, 4]] + + +def test_decode_masked_group_gemm_uses_primary_rows_only_when_eplb_is_enabled( + monkeypatch, +): + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True, num_logical_experts=8, world_size=1)) + captured = {} + + def masked(*args, **kwargs): + captured["w13"] = args[3] + captured["w13_scale"] = args[4] + captured["w2"] = args[5] + captured["w2_scale"] = args[6] + return "out" + + monkeypatch.setattr(deepgemm_module, "masked_group_gemm", masked) + pack = lambda: type( + "Pack", + (), + {"weight": torch.empty((10, 4)), "weight_scale": torch.empty((10, 1))}, + )() + + assert impl.masked_group_gemm((torch.empty((1, 4)),), pack(), pack(), torch.empty(8), torch.float16, 1) == "out" + assert captured["w13"].shape[0] == captured["w2"].shape[0] == 8 + assert captured["w13_scale"].shape[0] == captured["w2_scale"].shape[0] == 8 + + +def test_decode_fused_experts_uses_cached_primary_weight_packs_and_logical_experts( + monkeypatch, +): + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.n_routed_experts = 128 + _set_expert_parallel_state( + impl, + _test_parallel_state(eplb=True, num_logical_experts=128, world_size=16, num_redundant_experts_per_rank=2), + ) + impl.quant_method = object() + impl.ep_balance_counters = None + captured = [] + + def fused(**kwargs): + captured.append(kwargs) + return "out" + + monkeypatch.setattr(deepgemm_module, "fused_experts", fused) + pack = lambda: type( + "Pack", + (), + { + "weight": torch.empty((10, 4)), + "weight_scale": torch.empty((10, 1)), + "weight_zero_point": None, + }, + )() + w13, w2 = pack(), pack() + + for _ in range(2): + assert ( + impl._fused_experts( + torch.empty((1, 4)), + w13, + w2, + torch.ones((1, 2)), + torch.zeros((1, 2), dtype=torch.int64), + is_prefill=False, + ) + == "out" + ) + + assert [call["num_experts"] for call in captured] == [128, 128] + assert all(call["w13"].weight.shape[0] == call["w2"].weight.shape[0] == 8 for call in captured) + assert captured[0]["w13"] is captured[1]["w13"] + assert captured[0]["w2"] is captured[1]["w2"] + + +def test_transfer_plan_uses_existing_rows_and_prefers_local_node_replicas(): + current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) + target = current.clone() + target[0, 0] = 6 # primary r3, but r1 replica is on r0's node. + target[2, 1] = 4 # primary r2 is local to destination r2. + plan = build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) + by_dst = {(step.dst_rank, step.dst_slot): step for step in plan} + assert len(by_dst) == 2 + assert by_dst[0, 0] == TransferStep(0, 0, 1, 2) + assert by_dst[2, 1] == TransferStep(2, 1, 2, 0) + + +def test_transfer_plan_cross_node_and_stable_source_load_tie_break(): + current = torch.tensor([[0, 1], [2, 3], [4, 5], [4, 7]]) + target = current.clone() + target[0, 0] = 4 + target[0, 1] = 4 + first = build_transfer_plan(current, target, 8, 4, 2) + second = build_transfer_plan(current, target, 8, 4, 2) + assert first == second + selected = [step for step in first if step.dst_rank == 0] + assert [(step.src_rank, step.src_local_row) for step in selected] == [ + (2, 0), + (3, 2), + ] + + +def test_extract_expert_tensors_includes_weight_scale_and_zero_point_in_order(): + class Pack: + def __init__(self, offset, scale=True, zero=True): + self.weight = torch.full((3, 2), offset) + self.weight_scale = torch.full((3, 1), offset + 1) if scale else None + self.weight_zero_point = torch.full((3, 1), offset + 2) if zero else None + + weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False, zero=False)})() + tensors = extract_expert_tensors(weight) + assert [name for name, _ in tensors] == [ + "w13.weight", + "w13.weight_scale", + "w13.weight_zero_point", + "w2.weight", + ] + + +def test_commit_staging_rows_only_overwrites_redundant_rows(): + live = torch.arange(20).reshape(5, 4) + staging = torch.full((2, 4), -1) + commit_staging_rows( + live, + staging, + num_experts_per_rank=3, + changed_dst_slots=(0, 1), + ) + assert torch.equal(live[:3], torch.arange(12).reshape(3, 4)) + assert torch.equal(live[3:], staging) + + +def test_commit_staging_rows_preserves_unchanged_destination_slots(): + live = torch.arange(28).reshape(7, 4) + staging = torch.tensor([[-1, -1, -1, -1], [-2, -2, -2, -2], [-3, -3, -3, -3], [-4, -4, -4, -4]]) + original = live.clone() + + commit_staging_rows( + live, + staging, + num_experts_per_rank=3, + changed_dst_slots=(3, 1), + ) + + assert torch.equal(live[:4], original[:4]) + assert torch.equal(live[4], staging[1]) + assert torch.equal(live[5], original[5]) + assert torch.equal(live[6], staging[3]) + + +def test_commit_staging_rows_merges_contiguous_changed_slots(): + copies = [] + + class View: + def __init__(self, owner, start, length): + self.owner = owner + self.start = start + self.length = length + + def copy_(self, source, **_kwargs): + copies.append( + ( + self.owner, + self.start, + self.length, + source.owner, + source.start, + source.length, + ) + ) + + class Tensor: + def __init__(self, owner, rows): + self.owner = owner + self.shape = (rows,) + + def narrow(self, _dim, start, length): + return View(self.owner, start, length) + + commit_staging_rows( + Tensor("live", 20), + Tensor("staging", 4), + num_experts_per_rank=10, + changed_dst_slots=(3, 1, 2), + ) + + assert copies == [("live", 11, 3, "staging", 1, 3)] + + +def test_manager_inflight_ready_gate_commits_ordered_prefix_and_propagates_worker_error( + monkeypatch, +): + class Transfer: + def __init__(self): + self.pending = [(0, 0), (1, 1), (2, 2)] + self.commits = [] + self.finished = 0 + + def pending_layers(self): + return self.pending + + def commit(self, layer, buffer_index, post_copy=None): + assert self.pending.pop(0) == (layer, buffer_index) + self.commits.append((layer, buffer_index)) + if post_copy is not None: + post_copy() + + def finish(self): + self.finished += 1 + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.transfer = Transfer() + manager.in_flight = True + manager.world_size = 2 + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager.in_flight_layers = [0, 1, 2] + manager._commit_layer_metadata = lambda layer: committed.append(layer) + manager._finish_rebalance = lambda: finished.append(True) + committed, finished = [], [] + operations = [] + current_stream_calls = [] + + class CurrentStream: + def wait_stream(self, stream): + operations.append(("wait", stream)) + + overlap_stream = object() + + def current_stream(): + current_stream_calls.append(True) + return CurrentStream() + + monkeypatch.setattr(manager_module.torch.cuda, "current_stream", current_stream) + monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) + + original_commit = manager.transfer.commit + + def record_commit(*args, **kwargs): + operations.append(("commit", args[0])) + return original_commit(*args, **kwargs) + + manager.transfer.commit = record_commit + + def set_global_ready(count): + return lambda tensor, **kwargs: tensor.fill_(count) + + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(0)) + manager._poll_in_flight() + assert manager.transfer.commits == [] + assert operations == [] + assert current_stream_calls == [] + + # Local rank has three prefetched layers, but global MIN-ready only permits two. + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(2)) + manager._poll_in_flight() + assert manager.transfer.commits == [(0, 0), (1, 1)] + assert committed == [0, 1] + assert not finished + assert manager.transfer.finished == 0 + assert operations == [("wait", overlap_stream), ("commit", 0), ("commit", 1)] + assert current_stream_calls == [True] + + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) + manager._poll_in_flight() + assert finished == [True] + assert manager.transfer.finished == 1 + assert operations == [ + ("wait", overlap_stream), + ("commit", 0), + ("commit", 1), + ("wait", overlap_stream), + ("commit", 2), + ] + assert current_stream_calls == [True, True] + + manager.in_flight_layers = [3] + manager.transfer.pending = [(9, 0)] + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) + with pytest.raises(RuntimeError, match="does not match expected"): + manager._poll_in_flight() + + class BrokenTransfer: + def pending_layers(self): + raise RuntimeError("boom") + + manager.transfer = BrokenTransfer() + encoded_statuses = [] + + def retain_local_error(tensor, **_kwargs): + encoded_statuses.append(int(tensor.item())) + + monkeypatch.setattr(manager_module.dist, "all_reduce", retain_local_error) + with pytest.raises(RuntimeError, match="EPLB transfer worker failed on this rank") as exc_info: + manager._poll_in_flight() + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert str(exc_info.value.__cause__) == "boom" + assert encoded_statuses == [manager_module.EPLB_CONTROL_ERROR] + + +def test_manager_inflight_remote_worker_error_does_not_commit(monkeypatch): + class Transfer: + def __init__(self): + self.commits = [] + + def pending_layers(self): + return [(0, 0)] + + def commit(self, *args): + self.commits.append(args) + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.transfer = Transfer() + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager.in_flight_layers = [0] + manager._commit_layer_metadata = lambda _layer: None + manager._finish_rebalance = lambda: None + statuses = [] + + def remote_error(tensor, **_kwargs): + statuses.append(int(tensor.item())) + tensor.fill_(manager_module.EPLB_CONTROL_ERROR) + + monkeypatch.setattr(manager_module.dist, "all_reduce", remote_error) + + with pytest.raises(RuntimeError, match="EPLB transfer worker failed on another rank"): + manager._poll_in_flight() + assert statuses == [1] + assert manager.transfer.commits == [] + + +def test_evaluation_ready_gate_propagates_local_and_remote_errors(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager._evaluation_lock = threading.Lock() + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager._evaluation_error = RuntimeError("evaluation boom") + manager._evaluation_result = None + statuses = [] + + def retain_local_error(tensor, **_kwargs): + statuses.append(int(tensor.item())) + + monkeypatch.setattr(manager_module.dist, "all_reduce", retain_local_error) + with pytest.raises(RuntimeError, match="EPLB evaluation failed on this rank") as exc_info: + manager._evaluation_ready_on_all_ranks() + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert str(exc_info.value.__cause__) == "evaluation boom" + assert statuses == [manager_module.EPLB_CONTROL_ERROR] + + manager._evaluation_error = None + manager._evaluation_result = {"kind": "no_improvement"} + statuses.clear() + + def remote_error(tensor, **_kwargs): + statuses.append(int(tensor.item())) + tensor.fill_(manager_module.EPLB_CONTROL_ERROR) + + monkeypatch.setattr(manager_module.dist, "all_reduce", remote_error) + with pytest.raises(RuntimeError, match="EPLB evaluation failed on another rank"): + manager._evaluation_ready_on_all_ranks() + assert statuses == [1] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_manager_inflight_commit_orders_live_weights_between_overlap_forwards( + monkeypatch, +): + class Transfer: + def __init__(self, live, staging): + self.live = live + self.staging = staging + self.pending = [(0, 0)] + + def pending_layers(self): + return self.pending + + def commit(self, layer, buffer_index, post_copy=None): + assert self.pending.pop(0) == (layer, buffer_index) + self.live.copy_(self.staging, non_blocking=True) + if post_copy is not None: + post_copy() + + def finish(self): + pass + + live = torch.tensor([1.0], device="cuda") + staging = torch.tensor([2.0], device="cuda") + previous_read = torch.empty_like(live) + next_read = torch.empty_like(live) + source_stream = torch.cuda.Stream(device=live.device) + destination_stream = torch.cuda.Stream(device=live.device) + initial_stream = torch.cuda.current_stream(device=live.device) + original_overlap_stream = g_infer_context.overlap_stream + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.transfer = Transfer(live, staging) + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager.in_flight_layers = [0] + manager._commit_layer_metadata = lambda _layer: None + manager._finish_rebalance = lambda: None + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) + + try: + g_infer_context.overlap_stream = source_stream + with torch.cuda.stream(source_stream): + source_stream.wait_stream(initial_stream) + torch.cuda._sleep(20_000_000) + previous_read.copy_(live, non_blocking=True) + with torch.cuda.stream(destination_stream): + manager._poll_in_flight() + with torch.cuda.stream(source_stream): + source_stream.wait_stream(destination_stream) + next_read.copy_(live, non_blocking=True) + source_stream.synchronize() + + assert previous_read.item() == 1.0 + assert next_read.item() == 2.0 + finally: + g_infer_context.overlap_stream = original_overlap_stream + + +def test_transfer_ring_reuses_a_buffer_only_after_commit_and_consumption(monkeypatch): + operations = [] + + class Event: + def __init__(self): + self.recorded = 0 + self.synchronized = 0 + + def record(self, stream): + self.recorded += 1 + + def synchronize(self): + self.synchronized += 1 + operations.append("consumed synchronize") + + transfer = object.__new__(transfer_module._EPLBTransferBase) + transfer.backend = "test" + transfer.device = torch.device("cuda", 0) + transfer.staging_depth = 2 + transfer.staging = [[], []] + transfer.live = [[], [], []] + transfer.num_experts_per_rank = 0 + transfer._release = [threading.Event(), threading.Event()] + for release in transfer._release: + release.set() + transfer._consumed_events = [Event(), Event()] + transfer._consumed_recorded = [False, False] + transfer._changed_dst_slots = [(), ()] + transfer._pending = deque() + transfer._pending_lock = threading.Lock() + transfer._error = None + transfer._thread = None + transfer._needs_staging_reuse_barrier = True + transfer.transfer_group = "transfer-group" + copied = [] + + def copy_layer(layer, _plan, _staging): + copied.append(layer) + operations.append(("copy", layer)) + + transfer._copy_layer = copy_layer + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda device: None) + monkeypatch.setattr(transfer_module.torch.cuda, "Event", Event) + monkeypatch.setattr(transfer_module.torch.cuda, "current_stream", lambda: object()) + monkeypatch.setattr( + transfer_module.dist, + "barrier", + lambda **kwargs: operations.append(("barrier", kwargs["group"])), + ) + + transfer.start([(0, []), (1, []), (2, [])]) + deadline = time.monotonic() + 2 + while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: + time.sleep(0.001) + assert transfer.pending_layers() == [(0, 0), (1, 1)] + assert copied == [0, 1] + assert operations == [("copy", 0), ("copy", 1)] + + transfer.commit(0, 0) + while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: + time.sleep(0.001) + assert transfer.pending_layers() == [(1, 1), (2, 0)] + assert copied == [0, 1, 2] + assert transfer._consumed_events[0].synchronized == 1 + assert operations == [ + ("copy", 0), + ("copy", 1), + "consumed synchronize", + ("barrier", "transfer-group"), + ("copy", 2), + ] + transfer.commit(1, 1) + transfer.commit(2, 0) + transfer.finish() + + +def test_transfer_finalization_failure_stays_in_worker_and_success_finalizes_once(monkeypatch): + def make_transfer(finalize): + transfer = object.__new__(transfer_module._EPLBTransferBase) + transfer.backend = "test" + transfer.device = torch.device("cuda", 0) + transfer.global_rank = 0 + transfer.staging_depth = 1 + transfer.staging = [[]] + transfer._release = [threading.Event()] + transfer._release[0].set() + transfer._consumed_events = [object()] + transfer._consumed_recorded = [False] + transfer._changed_dst_slots = [()] + transfer._pending = deque() + transfer._pending_lock = threading.Lock() + transfer._error = None + transfer._thread = None + transfer._needs_staging_reuse_barrier = False + transfer._copy_batch = lambda _batch: None + transfer._finish_transfer_generation = finalize + return transfer + + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + + finalized_before_publish = [] + success = make_transfer(lambda: finalized_before_publish.append(len(success._pending))) + success.start([(0, [])]) + success.finish() + assert finalized_before_publish == [0] + assert success.pending_layers() == [(0, 0)] + + failed = make_transfer(lambda: (_ for _ in ()).throw(RuntimeError("cache boom"))) + failed.start([(0, [])]) + failed._thread.join() + assert list(failed._pending) == [] + with pytest.raises(RuntimeError, match="EPLB migration worker failed") as exc_info: + failed.pending_layers() + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert str(exc_info.value.__cause__) == "cache boom" + with pytest.raises(RuntimeError, match="EPLB migration worker failed"): + failed.finish() + assert failed._thread is None + + +def test_manager_rearms_after_rebalance_for_interval_one(): + recording_calls = [] + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.step_interval = 1 + manager.sampling_interval = 1 + manager._sampling_pending = False + manager._continuous_collection_start_step = 0 + manager.weights = [] + manager._eplb_states = [] + manager.target_placement = torch.zeros(1) + manager.in_flight_started_at = 0 + manager.global_rank = 1 + manager._set_recording = lambda enabled: recording_calls.append(enabled) + manager._finish_rebalance() + assert manager.in_flight is False + assert recording_calls == [True] + assert not manager._sampling_pending + assert manager._continuous_collection_start_step is None + + +def test_manager_sparse_insufficient_schedules_bounded_fresh_window(monkeypatch): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "insufficient", + "minimum_layer_samples": 1, + "minimum": 2, + } + manager.global_rank = 0 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.prefill_steps = 37 + manager.weights = [] + manager._eplb_states = [] + manager._continuous_collection_start_step = None + manager._continuous_collection_end_step = None + recordings, resets, logs = [], [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) + + assert not manager._poll_evaluation() + assert not hasattr(manager, "_retained_local_samples") + assert manager._continuous_collection_start_step == 40 + assert manager._continuous_collection_end_step == 60 + assert recordings == [False] + assert resets == [True] + assert manager.sampling_interval == 20 + assert "insufficient samples" in logs[0][0] + assert "scheduled_fresh_window" in logs[0][0] + + manager.in_flight = False + manager.evaluation_in_flight = False + starts = [] + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(manager.prefill_steps)) + manager.step() + manager.step() + assert manager.prefill_steps == 39 + assert starts == [] + manager.step() + assert manager.prefill_steps == 40 + assert starts == [] + for _ in range(19): + manager.step() + assert manager.prefill_steps == 59 + assert starts == [] + manager.step() + assert starts == [60] + + +def test_manager_full_window_insufficient_clears_and_backs_off(): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "insufficient", + "minimum_layer_samples": 1, + "minimum": 2, + } + manager.global_rank = 1 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.prefill_steps = 60 + manager.weights = [] + manager._eplb_states = [] + manager._continuous_collection_start_step = 40 + manager._continuous_collection_end_step = 60 + recordings, resets = [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + + assert not manager._poll_evaluation() + assert not hasattr(manager, "_retained_local_samples") + assert manager._continuous_collection_start_step is None + assert manager._continuous_collection_end_step is None + assert manager.sampling_interval == 80 + assert recordings == [False] + assert resets == [True] + + +def test_begin_continuous_collection_uses_full_window_at_fixed_boundary(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.evaluation_in_flight = False + manager.prefill_steps = 36 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager._sampling_pending = True + recordings, resets, starts = [], [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) + + # The window is never truncated to the next boundary: it waits until 40, + # records a full 20 fresh steps, then evaluates at the boundary at 60. + manager._begin_continuous_collection() + assert manager._continuous_collection_start_step == 40 + assert manager._continuous_collection_end_step == 60 + assert recordings == [False] + assert resets == [True] + assert not manager._sampling_pending + + for _ in range(4): + manager.step() + assert manager.prefill_steps == 40 + assert recordings == [False, True] + assert starts == [] + for _ in range(19): + manager.step() + assert starts == [] + manager.step() # 60: the full window ends and triggers the evaluation. + assert starts == [True] + + +def test_begin_continuous_collection_preserves_full_window_at_sparse_boundary(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.prefill_steps = 80 + manager.step_interval = 20 + manager.sampling_interval = 80 + manager._sampling_pending = False + recordings = [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: None + + manager._begin_continuous_collection() + assert manager._continuous_collection_start_step == 140 + assert manager._continuous_collection_end_step == 160 + assert recordings == [False] + + +def test_first_no_improvement_switches_to_sparse_sampling_window(monkeypatch): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 2, + } + manager.global_rank = 1 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager._continuous_collection_start_step = 0 + manager.weights = [] + manager._eplb_states = [] + recordings = [] + manager._set_recording = lambda enabled: recordings.append(enabled) + + assert not manager._poll_evaluation() + assert manager._continuous_collection_start_step is None + assert recordings == [False] + assert manager.sampling_interval == 80 + + +def test_continuous_collection_evaluates_only_after_one_full_base_window(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager._continuous_collection_start_step = 0 + manager._continuous_collection_end_step = 20 + manager.prefill_steps = 0 + manager.step_interval = 20 + manager.sampling_interval = 320 + manager._sampling_pending = False + manager.evaluation_in_flight = False + started = [] + monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) + + for _ in range(19): + manager.step() + assert manager.prefill_steps == 19 + assert started == [] + + manager.step() + assert manager.prefill_steps == 20 + assert started == [True] + + +def test_no_improvement_exponentially_backs_off_sampling_interval_at_cap(): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager._evaluation_lock = threading.Lock() + manager._set_recording = lambda _enabled: None + manager.global_rank = 1 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.weights = [] + manager._eplb_states = [] + + for expected_interval in (80, 320, 320): + manager.evaluation_in_flight = True + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 2, + } + assert not manager._poll_evaluation() + assert manager.sampling_interval == expected_interval + + +def test_sparse_backoff_arms_and_evaluates_only_at_new_interval_boundary(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.prefill_steps = 18 + manager.step_interval = 20 + manager.sampling_interval = 80 + manager._sampling_pending = False + manager._continuous_collection_start_step = None + manager._continuous_collection_end_step = None + manager.evaluation_in_flight = False + recordings, resets, starts = [], [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) + + manager.step() + manager.step() + assert manager.prefill_steps == 20 + assert recordings == [] + assert starts == [] + + for _ in range(55): + manager.step() + assert manager.prefill_steps == 75 + assert recordings == [] + assert starts == [] + + manager.step() + assert manager.prefill_steps == 76 + assert recordings == [True] + assert resets == [True] + assert manager._sampling_pending + + for _ in range(4): + manager.step() + assert manager.prefill_steps == 80 + assert starts == [True] + assert not manager._sampling_pending + + +def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.current_placement = torch.zeros((1, 1, 1), dtype=torch.int64) + manager.num_logical_experts = 1 + manager.world_size = 1 + manager.node_world_size = 1 + manager.step_interval = 20 + manager.sampling_interval = 320 + manager._continuous_collection_start_step = 0 + manager.global_rank = 1 + manager.transfer = type("Transfer", (), {"start": lambda self, plans: setattr(self, "plans", plans)})() + manager._reset_recorded_samples = lambda: None + + manager._start_rebalance( + { + "placement": torch.zeros((1, 1, 1), dtype=torch.int64), + "improved": torch.tensor([True]), + "metadata": [None], + "layer_plans": [(0, object())], + "before": {"max": 1.0, "p95": 1.0}, + "after": {"max": 1.0, "p95": 1.0}, + "model_imbalance_ratio": 1.0, + "candidate_model_imbalance_ratio": 1.0, + "candidate_rebalance_gain": 0.1, + "candidate_changed_layer_count": 1, + } + ) + + assert manager.sampling_interval == 20 + assert manager.in_flight + assert manager._continuous_collection_start_step is None + + +def test_first_rebalance_completion_switches_to_four_step_sparse_window(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.step_interval = 20 + manager.sampling_interval = manager.step_interval + manager._sampling_pending = False + manager.weights = [] + manager._eplb_states = [] + manager.target_placement = torch.zeros(1) + manager.in_flight_started_at = 0 + manager.global_rank = 1 + recordings, starts = [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._finish_rebalance() + assert recordings == [False] + + manager.in_flight = False + manager.prefill_steps = 38 + manager.evaluation_in_flight = False + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) + manager.prefill_steps = 35 + manager.step() + assert recordings == [False, True] + for _ in range(4): + manager.step() + assert starts == [True] + + +def test_manager_inflight_step_does_not_poll(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = True + calls = [] + manager._poll_in_flight = lambda: calls.append("poll") + manager.step() + assert calls == [] + + +def test_manager_poll_advances_inflight_before_evaluation(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = True + manager.evaluation_in_flight = True + calls = [] + manager._poll_in_flight = lambda: calls.append("inflight") + manager._poll_evaluation = lambda: calls.append("evaluation") + + manager.poll() + + assert calls == ["inflight"] + + +def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_result = None + manager._evaluation_error = None + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + calls = [] + manager._poll_evaluation = lambda: calls.append("evaluation") + + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(0)) + manager.poll() + assert calls == [] + + manager._evaluation_result = {"kind": "no_improvement"} + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) + manager.poll() + assert calls == ["evaluation"] + + +def test_nixl_descriptor_runs_merge_only_jointly_contiguous_source_and_destination_rows(): + steps = [ + TransferStep(0, 2, 1, 3), + TransferStep(0, 0, 1, 1), + TransferStep(0, 1, 1, 2), + TransferStep(0, 4, 1, 7), + ] + runs = transfer_module.NixlEPLBTransfer._contiguous_runs(steps) + assert [[(step.dst_slot, step.src_local_row) for step in run] for run in runs] == [ + [(0, 1), (1, 2), (2, 3)], + [(4, 7)], + ] + + +def test_nixl_remote_read_cache_reuses_exact_batch_key_and_releases_on_shutdown(): + class Tensor: + nbytes = 8 + + def __init__(self, pointer): + self.pointer = pointer + + def data_ptr(self): + return self.pointer + + def get_device(self): + return 0 + + def __getitem__(self, _): + return self + + class Agent: + def __init__(self): + self.prepared = 0 + self.made = 0 + self.released_xfers = 0 + self.released_dlists = 0 + self.removed_agents = [] + + def get_xfer_descs(self, descriptors, _): + return descriptors + + def prep_xfer_dlist(self, *_args, **_kwargs): + self.prepared += 1 + return f"dlist-{self.prepared}" + + def make_prepped_xfer(self, *_args, **_kwargs): + self.made += 1 + return f"xfer-{self.made}" + + def query_xfer_backend(self, _): + return "UCX" + + def release_xfer_handle(self, _): + self.released_xfers += 1 + + def release_dlist_handle(self, _): + self.released_dlists += 1 + + def remove_remote_agent(self, remote_name): + self.removed_agents.append(remote_name) + + agent = Agent() + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer._nixl_agent = agent + transfer._xfer_cache = {} + transfer._used_xfer_cache_keys = set() + transfer._remote_agents = {1: "remote-1"} + transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} + transfer._registered_descs = None + transfer.live = [[("weight", Tensor(100))]] + staging = [("weight", Tensor(200))] + first_entries = [(0, [TransferStep(0, 0, 1, 0)], staging)] + + first = transfer._get_remote_read(1, first_entries) + assert transfer._get_remote_read(1, first_entries) is first + assert agent.made == 1 + + changed_entries = [(0, [TransferStep(0, 1, 1, 0)], staging)] + transfer._get_remote_read(1, changed_entries) + assert agent.made == 2 + + transfer.shutdown() + assert agent.released_xfers == 2 + assert agent.released_dlists == 4 + assert agent.removed_agents == ["remote-1"] + + +def test_nixl_descriptor_caches_are_bounded_to_the_current_transfer_generation( + monkeypatch, +): + class Tensor: + nbytes = 8 + + def __init__(self, pointer): + self.pointer = pointer + + def data_ptr(self): + return self.pointer + + def get_device(self): + return 0 + + def __getitem__(self, _): + return self + + class Agent: + def __init__(self): + self.made = 0 + self.released_xfers = 0 + self.released_dlists = 0 + + def get_xfer_descs(self, descriptors, _): + return descriptors + + def prep_xfer_dlist(self, *_args, **_kwargs): + return object() + + def make_prepped_xfer(self, *_args, **_kwargs): + self.made += 1 + return object() + + def query_xfer_backend(self, _): + return "UCX" + + def release_xfer_handle(self, _): + self.released_xfers += 1 + + def release_dlist_handle(self, _): + self.released_dlists += 1 + + def remove_remote_agent(self, _): + pass + + agent = Agent() + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer._nixl_agent = agent + transfer._xfer_cache = {} + transfer._used_xfer_cache_keys = set() + transfer._push_descriptor_cache = {} + transfer._used_push_descriptor_cache_keys = set() + transfer._remote_agents = {1: "remote-1"} + transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} + transfer._registered_descs = None + transfer._ipc_staging = {} + transfer.live = [[("weight", Tensor(100))]] + transfer.device = torch.device("cuda", 0) + transfer.transfer_group = object() + transfer.global_rank = 0 + transfer.world_size = 2 + transfer.staging_depth = 1 + transfer.staging = [[]] + transfer.num_experts_per_rank = 0 + transfer._release = [threading.Event()] + transfer._release[0].set() + transfer._consumed_events = [object()] + transfer._consumed_recorded = [False] + transfer._changed_dst_slots = [()] + transfer._pending = deque() + transfer._pending_lock = threading.Lock() + transfer._error = None + transfer._thread = None + transfer._needs_staging_reuse_barrier = False + + staging = [("weight", Tensor(200))] + entries_a = [(0, [TransferStep(0, 0, 1, 0)], staging)] + entries_b = [(0, [TransferStep(0, 1, 1, 0)], staging)] + source_a, destination_a = Tensor(300), Tensor(400) + source_b, destination_b = Tensor(500), Tensor(600) + copies_a = [(destination_a, source_a)] + copies_b = [(destination_b, source_b)] + key_a = ((source_a.data_ptr(), destination_a.data_ptr()),) + key_b = ((source_b.data_ptr(), destination_b.data_ptr()),) + stale_key = ((700, 800),) + transfer._push_descriptor_cache = { + key_a: (object(), object()), + stale_key: (object(), object()), + } + generation = [entries_a, copies_a] + + def copy_batch(_batch): + transfer._get_remote_read(1, generation[0]) + transfer._cached_descriptor_tensors(generation[1]) + + transfer._copy_batch = copy_batch + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + + transfer.start([(0, [])]) + transfer.finish() + assert agent.made == 1 + assert agent.released_xfers == agent.released_dlists == 0 + assert set(transfer._push_descriptor_cache) == {key_a} + assert len(transfer._xfer_cache) == 1 + + # The real manager releases each staging buffer through commit(). This + # focused cache test has no commits, so model that hand-off before the + # next generation reuses buffer zero. + transfer._release[0].set() + transfer.start([(0, [])]) + transfer.finish() + assert agent.made == 1 + assert agent.released_xfers == agent.released_dlists == 0 + assert set(transfer._push_descriptor_cache) == {key_a} + assert len(transfer._xfer_cache) == 1 + + transfer._push_descriptor_cache[key_b] = (object(), object()) + generation[:] = [entries_b, copies_b] + transfer._release[0].set() + transfer.start([(0, [])]) + transfer.finish() + assert agent.made == 2 + assert agent.released_xfers == 1 + assert agent.released_dlists == 2 + assert set(transfer._push_descriptor_cache) == {key_b} + assert len(transfer._xfer_cache) == 1 + transfer.shutdown() + + +def test_nixl_ipc_metadata_exports_staging_per_local_target(monkeypatch): + from lightllm.server.router.model_infer.mode_backend.pd import p2p_fix + + class Tensor: + shape = (4, 2) + dtype = torch.float16 + device = torch.device("cuda", 0) + nbytes = 16 + + def __init__(self, label): + self.label = label + + def numel(self): + return 3 + + def __getitem__(self, _index): + return self + + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer.transfer_group = object() + transfer.global_rank = 0 + transfer.world_size = 3 + transfer.device = torch.device("cuda", 0) + transfer.live = [[("w13.weight", Tensor("local-w13")), ("w2.weight", Tensor("local-w2"))]] + transfer.staging_depth = 8 + transfer.staging = [[("w13.weight", Tensor(f"local-staging-{index}"))] for index in range(8)] + transfer._ipc_staging = {} + + reduce_calls, rebuild_calls, gathers = [], [], [] + + def reduce_tensor(tensor): + reduce_calls.append(tensor.label) + return None, (f"export-{tensor.label}",) + + def rebuild_tensor(export): + rebuild_calls.append(export) + return Tensor(export) + + source_one = [[("w13.weight", (4, 2), torch.float16, (f"rank1-staging-{index}",))] for index in range(8)] + + def all_gather(output, value, **_kwargs): + gathers.append(value) + if len(gathers) == 1: + output[:] = ["node-a", "node-a", "node-b"] + else: + output[:] = [value, {0: {"staging": source_one}}, {}] + + monkeypatch.setattr(p2p_fix, "reduce_tensor", reduce_tensor) + monkeypatch.setattr(p2p_fix, "p2p_fix_rebuild_cuda_tensor", rebuild_tensor) + monkeypatch.setattr(transfer_module.dist, "all_gather_object", all_gather) + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + + transfer._init_ipc_metadata() + + assert reduce_calls == [*(f"local-staging-{index}" for index in range(8))] + assert rebuild_calls == [*(f"rank1-staging-{index}" for index in range(8))] + assert transfer._same_node_ranks == {0, 1} + assert transfer._cross_node_ranks == {2} + assert transfer._needs_staging_reuse_barrier + assert [name for name, _ in transfer._ipc_staging[1][0]] == ["w13.weight"] + + +def test_nixl_copy_batch_source_pushes_local_rows_and_keeps_remote_ucx_reads( + monkeypatch, +): + class Stream: + def __init__(self): + self.synchronized = 0 + + def synchronize(self): + self.synchronized += 1 + + @contextmanager + def use_stream(_stream): + yield + + stream = Stream() + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer.global_rank = 0 + transfer.num_experts_per_rank = 2 + transfer._same_node_ranks = {0, 1} + transfer._push_stream = stream + remote_reads, pushed, waited_xfers = [], [], [] + transfer._push_same_node = lambda dst_rank, entries: pushed.append((dst_rank, entries)) + transfer._get_remote_read = lambda src_rank, entries: remote_reads.append((src_rank, entries)) or ( + None, + None, + "xfer", + ) + transfer._wait_xfers = waited_xfers.extend + + monkeypatch.setattr(transfer_module.torch.cuda, "stream", use_stream) + staging = [] + local_step = TransferStep(1, 0, 0, 0) + remote_step = TransferStep(0, 1, 2, 0) + + transfer._copy_batch([(0, [local_step], 0, staging)]) + assert pushed == [(1, [(0, [local_step], 0)])] + assert remote_reads == [] + + transfer._copy_batch([(0, [remote_step], 0, staging)]) + assert [rank for rank, _ in remote_reads] == [2] + assert waited_xfers == [(None, None, "xfer")] + + +def test_nixl_copy_batch_self_only_rank_pushes_and_synchronizes(monkeypatch): + class Stream: + def __init__(self): + self.synchronized = 0 + + def synchronize(self): + self.synchronized += 1 + + @contextmanager + def use_stream(_stream): + yield + + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer.global_rank = 0 + transfer.num_experts_per_rank = 2 + transfer._same_node_ranks = {0} + transfer._push_stream = Stream() + pushed = [] + transfer._push_same_node = lambda dst_rank, entries: pushed.append((dst_rank, entries)) + transfer._wait_xfers = lambda _xfers: None + monkeypatch.setattr(transfer_module.torch.cuda, "stream", use_stream) + + self_step = TransferStep(0, 0, 0, 0) + transfer._copy_batch([(0, [self_step], 0, [])]) + + assert pushed == [(0, [(0, [self_step], 0)])] + assert transfer._push_stream.synchronized == 1 + + +def test_manager_constructs_nixl_transfer(monkeypatch): + weight = type( + "Weight", + (), + { + "n_routed_experts": 4, + "expert_parallel_state": _test_parallel_state( + eplb=True, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=2, + route_counter=torch.zeros((2, 4), dtype=torch.int64), + ), + }, + )() + transfer = object() + groups = [object(), object(), object()] + new_group_calls = [] + monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) + monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(manager_module, "get_node_world_size", lambda: 2) + monkeypatch.setattr(manager_module, "get_prefill_eplb_step_interval", lambda: 20) + monkeypatch.setattr(manager_module, "get_eplb_rebalance_gain_threshold", lambda: 0.07) + + def new_group(*args, **kwargs): + new_group_calls.append((args, kwargs)) + return groups[len(new_group_calls) - 1] + + monkeypatch.setattr(manager_module.dist, "new_group", new_group) + transfer_calls = [] + monkeypatch.setattr( + manager_module, + "NixlEPLBTransfer", + lambda weights, group, rank, world_size: ( + transfer_calls.append((weights, group, rank, world_size)) or transfer + ), + ) + logs = [] + monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) + manager = manager_module.EPLBManager(type("Model", (), {})()) + assert manager.transfer is transfer + assert ( + manager.evaluation_group, + manager.control_group, + manager.transfer_group, + ) == tuple(groups) + assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 + assert transfer_calls == [([weight], groups[2], 0, 2)] + assert manager.rebalance_gain_threshold == 0.07 + assert "rebalance_gain_threshold=0.0700" in logs[0] + assert manager._continuous_collection_start_step is None + assert manager._continuous_collection_end_step == manager.step_interval + assert weight.expert_parallel_state.eplb.recording + assert manager._eplb_states[0] is weight.expert_parallel_state.eplb + assert not hasattr(weight.expert_parallel_state.eplb, "record_load") + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +@pytest.mark.parametrize("record_load", [False, True]) +@pytest.mark.parametrize("tokens", [1, 32]) +@pytest.mark.parametrize("scoring_func", ["sigmoid", "softmax"]) +@pytest.mark.parametrize("renormalize", [False, True]) +def test_grouped_topk_eplb_matches_topk_mapping_and_counting(record_load, tokens, scoring_func, renormalize): + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import ( + triton_grouped_topk, + triton_grouped_topk_eplb, + ) + + torch.manual_seed(1234) + topk = 8 + experts = 256 + num_expert_group = 8 + gating_output = torch.randn((tokens, experts), dtype=torch.bfloat16, device="cuda") + correction_bias = torch.randn((experts,), dtype=torch.float32, device="cuda") + hidden_states = torch.empty((tokens, 1), dtype=torch.bfloat16, device="cuda") + logical_to_physical = torch.stack( + ( + torch.arange(experts, dtype=torch.int32, device="cuda"), + torch.arange(experts, dtype=torch.int32, device="cuda") + experts, + ), + dim=1, + ) + logical_replica_count = torch.where( + torch.arange(experts, device="cuda") % 3 == 0, + torch.full((experts,), 2, dtype=torch.int32, device="cuda"), + torch.ones((experts,), dtype=torch.int32, device="cuda"), + ) + expected_counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") + fused_counter = torch.zeros_like(expected_counter) + + expected_weights, logical_ids = triton_grouped_topk( + hidden_states, + gating_output, + correction_bias, + topk, + renormalize, + num_expert_group, + 4, + scoring_func, + 2, + ) + if tokens == 1: + replica_indices = torch.zeros_like(logical_ids) + else: + token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) + replica_indices = ( + (((token_indices * 2654435769) & 0xFFFFFFFF) + ((logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF)) + & 0xFFFFFFFF + ) % logical_replica_count[logical_ids.to(torch.long)].to(torch.int64) + expected_ids = logical_to_physical[logical_ids.to(torch.long), replica_indices.to(torch.long)].to(torch.long) + if record_load: + expected_counter[1].scatter_add_( + 0, + logical_ids.reshape(-1).to(torch.long), + torch.ones(logical_ids.numel(), dtype=torch.int64, device="cuda"), + ) + fused_weights, fused_ids, fused_logical_ids = triton_grouped_topk_eplb( + hidden_states, + gating_output, + correction_bias, + topk, + renormalize, + num_expert_group, + 4, + scoring_func, + logical_to_physical, + logical_replica_count, + fused_counter, + sample_index=1, + record_load=record_load, + use_grouped_topk=True, + group_score_used_topk_num=2, + ) + torch.cuda.synchronize() + + torch.testing.assert_close(fused_weights, expected_weights, rtol=1e-5, atol=1e-6) + assert torch.equal(fused_ids, expected_ids) + assert fused_logical_ids is None + assert torch.equal(fused_counter, expected_counter) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +@pytest.mark.parametrize("record_load", [False, True]) +@pytest.mark.parametrize("tokens", [1, 32]) +def test_global_topk_eplb_supports_logical_ids_and_counting(record_load, tokens): + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb + + torch.manual_seed(1234) + topk = 4 + experts = 64 + gating_output = torch.randn((tokens, experts), dtype=torch.float32, device="cuda") + logical_to_physical = torch.stack( + ( + torch.arange(experts, dtype=torch.int32, device="cuda"), + torch.arange(experts, dtype=torch.int32, device="cuda") + experts, + ), + dim=1, + ) + logical_replica_count = torch.where( + torch.arange(experts, device="cuda") % 3 == 0, + torch.full((experts,), 2, dtype=torch.int32, device="cuda"), + torch.ones((experts,), dtype=torch.int32, device="cuda"), + ) + expected_counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") + fused_counter = torch.zeros_like(expected_counter) + expected_weights, expected_logical_ids = torch.softmax(gating_output, dim=-1).topk(topk, dim=-1) + if tokens == 1: + replica_indices = torch.zeros_like(expected_logical_ids) + else: + token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) + replica_indices = ( + ( + ((token_indices * 2654435769) & 0xFFFFFFFF) + + ((expected_logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF) + ) + & 0xFFFFFFFF + ) % logical_replica_count[expected_logical_ids].to(torch.int64) + expected_ids = logical_to_physical[expected_logical_ids, replica_indices.to(torch.long)].to(torch.long) + if record_load: + expected_counter[1].scatter_add_( + 0, + expected_logical_ids.reshape(-1), + torch.ones(expected_logical_ids.numel(), dtype=torch.int64, device="cuda"), + ) + expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True) + + weights, physical_ids, logical_ids = triton_grouped_topk_eplb( + hidden_states=torch.empty((tokens, 1), dtype=torch.float32, device="cuda"), + gating_output=gating_output, + correction_bias=torch.randn((experts,), dtype=torch.float32, device="cuda"), + topk=topk, + renormalize=True, + num_expert_group=8, + topk_group=4, + scoring_func="sigmoid", + logical_to_physical_map=logical_to_physical, + logical_replica_count=logical_replica_count, + expert_counter=fused_counter, + sample_index=1, + record_load=record_load, + use_grouped_topk=False, + return_logical_ids=True, + ) + torch.cuda.synchronize() + + torch.testing.assert_close(weights, expected_weights, rtol=1e-5, atol=1e-6) + assert torch.equal(physical_ids, expected_ids) + assert torch.equal(logical_ids, expected_logical_ids) + assert torch.equal(fused_counter, expected_counter) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +def test_triton_grouped_topk_eplb_empty_tokens_skips_kernel(): + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb + + experts = 64 + counter = torch.zeros((1, experts), dtype=torch.int64, device="cuda") + weights, physical_ids, logical_ids = triton_grouped_topk_eplb( + hidden_states=torch.empty((0, 1), device="cuda"), + gating_output=torch.empty((0, experts), device="cuda"), + correction_bias=None, + topk=4, + renormalize=False, + num_expert_group=8, + topk_group=4, + scoring_func="softmax", + logical_to_physical_map=torch.zeros((experts, 1), dtype=torch.int32, device="cuda"), + logical_replica_count=torch.ones((experts,), dtype=torch.int32, device="cuda"), + expert_counter=counter, + sample_index=0, + record_load=True, + use_grouped_topk=False, + return_logical_ids=True, + ) + + assert weights.shape == physical_ids.shape == logical_ids.shape == (0, 4) + assert physical_ids.dtype is logical_ids.dtype is torch.long + assert torch.equal(counter, torch.zeros_like(counter)) diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py new file mode 100644 index 0000000000..4397943be4 --- /dev/null +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -0,0 +1,225 @@ +"""Two-GPU NIXL EPLB correctness and 512 MiB micro-performance test.""" +import os +import random +import socket +import statistics +import time + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + NixlEPLBTransfer, + build_transfer_plan, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_initial_redundant_expert_ids, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + EPLBState, + ExpertParallelState, +) + +pytest.importorskip("nixl", reason="NIXL package is required") + + +class _Pack: + def __init__(self, weight, weight_scale): + self.weight = weight + self.weight_scale = weight_scale + self.weight_zero_point = None + + +def _free_port(): + sock = socket.socket() + sock.bind(("127.0.0.1", 0)) + port = sock.getsockname()[1] + sock.close() + return port + + +class _FakeWeight: + def __init__(self, rank, layer_index, row_elements): + self.n_routed_experts = 32 + self.expert_parallel_state = ExpertParallelState( + num_logical_experts=32, + world_size=2, + eplb=EPLBState( + num_redundant_experts_per_rank=16, + initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids(32, 2, 16), + logical_to_physical_map=torch.zeros((32, 2), dtype=torch.int32, device="cuda"), + logical_replica_count=torch.ones(32, dtype=torch.int32, device="cuda"), + route_counter=torch.zeros((1, 32), dtype=torch.int64, device="cuda"), + ), + ) + base = rank * 100 + layer_index * 100 + self.w13 = self._pack(base, row_elements) + self.w2 = self._pack(base + 10, row_elements) + + @staticmethod + def _pack(base, row_elements): + weight = torch.empty((32, row_elements), dtype=torch.float16, device="cuda") + for row in range(weight.shape[0]): + weight[row].fill_(base + row) + scale = torch.empty((32, 1), dtype=torch.float32, device="cuda") + for row in range(scale.shape[0]): + scale[row].fill_(base + row + 0.5) + return _Pack(weight, scale) + + +def _wait_for_ready_prefix(transfer, control_group): + deadline = time.monotonic() + 30 + while True: + pending = transfer.pending_layers() + ready_count = torch.tensor([len(pending)], dtype=torch.int32) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=control_group) + if int(ready_count.item()) > 0: + return pending[: int(ready_count.item())] + if time.monotonic() >= deadline: + raise TimeoutError("EPLB transfer worker did not publish a globally ready layer") + time.sleep(0.001) + + +def _run_layers( + transfer, + control_group, + layer_plans, + callback=lambda layer_index: None, + before_commit_callback=lambda layer_index: None, +): + transfer.start(layer_plans) + committed = 0 + while committed < len(layer_plans): + pending = _wait_for_ready_prefix(transfer, control_group) + assert len(pending) <= len(layer_plans) - committed + for layer_index, buffer_index in pending: + assert layer_index == layer_plans[committed][0] + before_commit_callback(layer_index) + transfer.commit( + layer_index, + buffer_index, + lambda layer_index=layer_index: callback(layer_index), + ) + committed += 1 + transfer.finish() + + +def _assert_correctness(weights, rank, source_rows_by_dst_slot): + if rank == 0: + for layer_index in (0, len(weights) - 1): + base = 100 + layer_index * 100 + for dst_slot, src_row in enumerate(source_rows_by_dst_slot): + dst_row = 16 + dst_slot + assert torch.all(weights[layer_index].w13.weight[dst_row] == base + src_row) + assert torch.all(weights[layer_index].w13.weight_scale[dst_row] == base + src_row + 0.5) + assert torch.all(weights[layer_index].w2.weight[dst_row] == base + src_row + 10) + assert torch.all(weights[layer_index].w2.weight_scale[dst_row] == base + src_row + 10.5) + + +def _benchmark(transfer, control_group, layer_plans, payload): + for _ in range(3): + _run_layers(transfer, control_group, layer_plans) + dist.barrier(group=control_group) + started = time.perf_counter() + for _ in range(8): + _run_layers(transfer, control_group, layer_plans) + torch.cuda.synchronize() + return payload * 8 / (time.perf_counter() - started) / 1e9 + + +def _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload): + batch = [ + (layer_index, plan, buffer_index, transfer.staging[buffer_index]) + for buffer_index, (layer_index, plan) in enumerate(layer_plans) + ] + for _ in range(3): + transfer._copy_batch(batch) + dist.barrier(group=control_group) + samples = [] + for _ in range(20): + started = time.perf_counter() + transfer._copy_batch(batch) + samples.append(payload / (time.perf_counter() - started) / 1e9) + torch.cuda.synchronize() + median = statistics.median(samples) + print( + f"NIXL _copy_batch payload={payload / 2**20:.1f} MiB; " + f"min={min(samples):.2f} GB/s median={median:.2f} " + f"mean={statistics.mean(samples):.2f} max={max(samples):.2f}", + flush=True, + ) + return median + + +def _eplb_worker(rank, port, queue): + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = str(port) + torch.cuda.set_device(rank) + dist.init_process_group("gloo", rank=rank, world_size=2) + control_group = dist.new_group([0, 1], backend="gloo") + transfer_group = dist.new_group([0, 1], backend="gloo") + # Eight layers × 16 changed experts × two 2 MiB rows = 512 MiB useful remote weight payload. + row_elements = int(os.getenv("LIGHTLLM_EPLB_TEST_ROW_ELEMENTS", str(1024 * 1024))) + current = torch.tensor([list(range(16)), list(range(16))]) + source_rows_by_dst_slot = list(range(16)) + random.Random(20260731).shuffle(source_rows_by_dst_slot) + # Rank 0 receives the reverse logical-expert range in a deterministic random slot order. + # Consequently every descriptor has a distinct source and destination row. + target = torch.tensor([[16 + source_row for source_row in source_rows_by_dst_slot], list(range(16))]) + plan = build_transfer_plan(current, target, 32, 2, 2) + assert [step.src_local_row for step in plan if step.dst_rank == 0] == source_rows_by_dst_slot + benchmark_layer_count = 8 + layer_count = benchmark_layer_count + 1 + row_payload = ( + 2 * row_elements * torch.empty((), dtype=torch.float16).element_size() + + 2 * torch.empty((), dtype=torch.float32).element_size() + ) + payload = benchmark_layer_count * 16 * row_payload + weights = [_FakeWeight(rank, layer_index, row_elements) for layer_index in range(layer_count)] + transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=2) + assert transfer.staging_depth == 8 + assert transfer._eplb_states[0] is weights[0].expert_parallel_state.eplb + assert all(tensor.is_cuda for staging in transfer.staging for _, tensor in staging) + + wrap_layer_plans = [(layer_index, plan) for layer_index in range(layer_count)] + + def delay_rank_zero_first_commit(layer_index): + if rank == 0 and layer_index == 0: + time.sleep(0.1) + + _run_layers( + transfer, + control_group, + wrap_layer_plans, + before_commit_callback=delay_rank_zero_first_commit, + ) + torch.cuda.synchronize() + _assert_correctness(weights, rank, source_rows_by_dst_slot) + layer_plans = wrap_layer_plans[:benchmark_layer_count] + nixl_copy_batch = _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload) + nixl_bandwidth = _benchmark(transfer, control_group, layer_plans, payload) + transfer.shutdown() + + gathered = [None, None] + dist.all_gather_object(gathered, (nixl_bandwidth, nixl_copy_batch), group=control_group) + if rank == 0: + queue.put((payload, *gathered[0])) + dist.barrier(group=control_group) + dist.destroy_process_group() + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < 2, + reason="requires two CUDA GPUs", +) +def test_eplb_transfer_two_gpu_correctness_and_microperf(): + queue = mp.get_context("spawn").SimpleQueue() + mp.spawn(_eplb_worker, args=(_free_port(), queue), nprocs=2, join=True) + payload, nixl_gbps, nixl_copy_batch_gbps = queue.get() + print( + f"EPLB remote payload/round: {payload / 2**20:.1f} MiB; " + f"NIXL={nixl_gbps:.2f} GB/s NIXL _copy_batch={nixl_copy_batch_gbps:.2f} GB/s" + ) + assert nixl_gbps > 0 diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py new file mode 100644 index 0000000000..96028553ea --- /dev/null +++ b/unit_tests/server/test_api_start_eplb.py @@ -0,0 +1,28 @@ +import pytest + +from lightllm.server import api_start +from lightllm.server.core.objs.start_args_type import StartArgs + + +def test_eplb_prefill_cudagraph_is_rejected_before_starting_subprocesses(monkeypatch): + args = StartArgs( + enable_ep_moe=True, + enable_prefill_eplb=True, + enable_prefill_cudagraph=True, + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) + monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) + monkeypatch.setattr( + api_start.process_manager, + "start_submodule_processes", + lambda *args, **kwargs: pytest.fail("subprocess startup must not be reached"), + ) + + with pytest.raises(AssertionError, match="--enable_prefill_eplb does not support --enable_prefill_cudagraph"): + api_start._launch_subprocesses(args) From e6a07199907fd1d6f61941996b6bdb2511266ae8 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 2 Sep 2026 19:06:46 +0800 Subject: [PATCH 04/72] feat: support eplb for mtp --- .../fused_moe/expert_parallel_state.py | 20 ++++- .../fused_moe/fused_moe_weight.py | 3 +- lightllm/server/api_start.py | 1 - .../model_infer/mode_backend/base_backend.py | 7 +- unit_tests/common/fused_moe/test_eplb.py | 90 +++++++++---------- unit_tests/server/test_api_start_eplb.py | 36 ++++++++ 6 files changed, 107 insertions(+), 50 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py index fc0f11e015..f66d31a62b 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py @@ -1,9 +1,27 @@ +from contextlib import contextmanager +from contextvars import ContextVar from dataclasses import dataclass -from typing import Optional +from typing import Iterator, Optional import torch +_eplb_model_init_disabled: ContextVar[bool] = ContextVar("eplb_model_init_disabled", default=False) + + +def is_eplb_model_init_disabled() -> bool: + return _eplb_model_init_disabled.get() + + +@contextmanager +def disable_eplb_model_init() -> Iterator[None]: + token = _eplb_model_init_disabled.set(True) + try: + yield + finally: + _eplb_model_init_disabled.reset(token) + + @dataclass class EPLBState: num_redundant_experts_per_rank: int 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 225be834b2..0cd5a67522 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 @@ -12,6 +12,7 @@ from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( EPLBState, ExpertParallelState, + is_eplb_model_init_disabled, ) from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( @@ -92,7 +93,7 @@ def _init_expert_parallel_state(self): self._initial_redundant_expert_ids = [] self._initial_redundant_expert_idx_to_local_idx = {} eplb = None - if args.enable_prefill_eplb: + if args.enable_prefill_eplb and not is_eplb_model_init_disabled(): num_redundant_experts_per_rank = args.eplb_num_redundant_experts_per_rank all_initial_ids = build_initial_redundant_expert_ids( self.n_routed_experts, diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 31e5f39dbc..34ed299040 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -167,7 +167,6 @@ def _launch_subprocesses(args: StartArgs): assert ( args.eplb_num_redundant_experts_per_rank > 0 ), "--eplb_num_redundant_experts_per_rank must be greater than 0" - assert args.mtp_mode is None, "--enable_prefill_eplb does not support MTP modes" if args.enable_ep_moe: allowed_ep_prefill_att_backends = {"auto", "fa3", "triton", "flashqla"} diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 3291a17d0f..06b02b35ef 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -50,6 +50,9 @@ ) from .multi_level_kv_cache import MultiLevelKvCacheModule from lightllm.utils.profiler import ProcessProfiler, ProfilerCmd +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + disable_eplb_model_init, +) class ModeBackend: @@ -350,7 +353,9 @@ def init_mtp_draft_model(self, main_kvargs: dict): model_cfg=draft_model_cfg, spec_mode=spec_mode, ) - self.draft_models.append(draft_model_class(draft_model_kvargs)) + with disable_eplb_model_init(): + draft_model = draft_model_class(draft_model_kvargs) + self.draft_models.append(draft_model) self.logger.info(f"loaded speculative draft model class {self.draft_models[i].__class__}") return diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 38de738a48..6105cd0724 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -42,6 +42,10 @@ from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe import ( fused_moe_weight as fused_weight_module, ) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + disable_eplb_model_init, + is_eplb_model_init_disabled, +) from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( TransferStep, align_target_placement, @@ -49,7 +53,6 @@ commit_staging_rows, extract_expert_tensors, ) -from lightllm.utils import envs_utils def _test_parallel_state( @@ -262,51 +265,6 @@ def __init__(self, layer_num, enable_ep_moe=True): assert manager_module._find_fused_moe_weights(model) == [alternate, aliased, first] -def test_get_eplb_rebalance_gain_threshold_defaults_to_five_percent(monkeypatch): - monkeypatch.delenv("LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD", raising=False) - envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() - - try: - assert envs_utils.get_eplb_rebalance_gain_threshold() == 0.05 - finally: - envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() - - -def test_get_eplb_placement_stickiness_defaults_to_ten_percent(monkeypatch): - monkeypatch.delenv("LIGHTLLM_EPLB_PLACEMENT_STICKINESS", raising=False) - envs_utils.get_eplb_placement_stickiness.cache_clear() - - try: - assert envs_utils.get_eplb_placement_stickiness() == 0.1 - finally: - envs_utils.get_eplb_placement_stickiness.cache_clear() - - -@pytest.mark.parametrize(("configured", "expected"), [("0", 0.0), (".04", 0.04), ("1", 1.0)]) -def test_get_eplb_rebalance_gain_threshold_reads_valid_values(monkeypatch, configured, expected): - monkeypatch.setenv("LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD", configured) - envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() - - try: - assert envs_utils.get_eplb_rebalance_gain_threshold() == expected - finally: - envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() - - -@pytest.mark.parametrize("configured", ["-0.01", "1.01", "nan", "inf"]) -def test_get_eplb_rebalance_gain_threshold_rejects_invalid_values(monkeypatch, configured): - env_name = "LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD" - monkeypatch.setenv(env_name, configured) - envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() - - try: - with pytest.raises(ValueError, match=env_name): - envs_utils.get_eplb_rebalance_gain_threshold() - finally: - envs_utils.get_eplb_rebalance_gain_threshold.cache_clear() - monkeypatch.delenv(env_name, raising=False) - - def test_eplb_redundant_experts_defaults_per_ep_rank(): parser = make_argument_parser() @@ -1444,6 +1402,46 @@ def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): assert not hasattr(weight, "routed_expert_counter_tensor") +def test_disable_eplb_model_init_skips_eplb_state(monkeypatch): + args = type( + "Args", + (), + {"enable_prefill_eplb": True, "eplb_num_redundant_experts_per_rank": 2}, + )() + weight = object.__new__(fused_weight_module.FusedMoeWeight) + weight.n_routed_experts = 4 + weight.global_world_size = 2 + weight.global_rank_ = 0 + weight.enable_ep_moe = True + monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) + monkeypatch.setattr( + fused_weight_module, + "build_initial_redundant_expert_ids", + lambda *args, **kwargs: pytest.fail("disabled scope must not initialize EPLB"), + ) + + with disable_eplb_model_init(): + weight._init_expert_parallel_state() + + assert weight.expert_parallel_state.eplb is None + assert weight.expert_parallel_state.num_total_physical_experts == weight.n_routed_experts + assert weight._initial_redundant_expert_ids == [] + + +def test_disable_eplb_model_init_scope_restores_after_exception(): + assert not is_eplb_model_init_disabled() + with disable_eplb_model_init(): + assert is_eplb_model_init_disabled() + with disable_eplb_model_init(): + assert is_eplb_model_init_disabled() + assert is_eplb_model_init_disabled() + + with pytest.raises(RuntimeError): + with disable_eplb_model_init(): + raise RuntimeError + assert not is_eplb_model_init_disabled() + + def test_manager_evaluation_collective_preserves_source_node_axis(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.weights = [ diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py index 96028553ea..8064ab7952 100644 --- a/unit_tests/server/test_api_start_eplb.py +++ b/unit_tests/server/test_api_start_eplb.py @@ -26,3 +26,39 @@ def test_eplb_prefill_cudagraph_is_rejected_before_starting_subprocesses(monkeyp with pytest.raises(AssertionError, match="--enable_prefill_eplb does not support --enable_prefill_cudagraph"): api_start._launch_subprocesses(args) + + +def test_eplb_mtp_combination_is_not_rejected_before_starting_subprocesses(monkeypatch): + args = StartArgs( + model_dir="test-model", + enable_ep_moe=True, + enable_prefill_eplb=True, + mtp_mode="vanilla_no_att", + mtp_step=1, + eos_id=0, + data_type="float16", + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) + monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) + monkeypatch.setattr(api_start, "is_sm100_gpu", lambda: False) + monkeypatch.setattr(api_start, "auto_set_response_parsers", lambda args: None) + monkeypatch.setattr(api_start, "auto_configure_allreduce_flags_from_args", lambda args: None) + monkeypatch.setattr(api_start, "validate_ports", lambda ports: None) + monkeypatch.setattr(api_start, "set_env_start_args", lambda args: None) + monkeypatch.setattr(api_start, "get_shm_port_args", lambda create=False: None) + monkeypatch.setattr(api_start, "send_and_receive_node_ip", lambda args: None) + monkeypatch.setattr( + api_start.process_manager, + "start_submodule_processes", + lambda *args, **kwargs: (object(), None), + ) + monkeypatch.setattr(api_start.process_manager, "setup_exit_controller", lambda: None) + monkeypatch.setattr(api_start.process_manager, "register_process_tree", lambda process: None) + + api_start._launch_subprocesses(args) From 14af0f763006e839b37a62d48d801cacc22585ca Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Fri, 4 Sep 2026 02:54:50 +0000 Subject: [PATCH 05/72] perf: use CUDA memcpy for EPLB expert transfers --- .../triton_kernel/fused_moe/eplb_kernels.py | 55 --- .../triton_kernel/fused_moe/grouped_topk.py | 10 +- .../model_infer/mode_backend/eplb_manager.py | 3 +- .../model_infer/mode_backend/eplb_transfer.py | 394 ++++++++++------ unit_tests/common/fused_moe/test_eplb.py | 432 +++++++++++++++--- .../fused_moe/test_eplb_transfer_gpu.py | 225 ++++++++- 6 files changed, 857 insertions(+), 262 deletions(-) delete mode 100644 lightllm/common/basemodel/triton_kernel/fused_moe/eplb_kernels.py diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_kernels.py b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_kernels.py deleted file mode 100644 index 15462dff0b..0000000000 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_kernels.py +++ /dev/null @@ -1,55 +0,0 @@ -import torch -import triton -import triton.language as tl - - -@triton.jit -def eplb_replica_index(token_index, logical_id, replica_count): - """Choose a replica with independent phases for a token's top-k experts.""" - token_hash = token_index.to(tl.uint32) * 2654435769 - expert_hash = logical_id.to(tl.uint32) * 2246822519 - return (token_hash + expert_hash) % replica_count.to(tl.uint32) - - -@triton.jit -def _eplb_push_copy_kernel( - src_ptrs_ptr, - dst_ptrs_ptr, - bytes_per_descriptor, - BLOCK_SIZE: tl.constexpr, - ITEMS_PER_PROGRAM: tl.constexpr, -): - descriptor_index = tl.program_id(1) - offsets = tl.program_id(0) * (BLOCK_SIZE * ITEMS_PER_PROGRAM) + tl.arange(0, BLOCK_SIZE) - src_ptr = tl.load(src_ptrs_ptr + descriptor_index).to(tl.pointer_type(tl.uint64)) - dst_ptr = tl.load(dst_ptrs_ptr + descriptor_index).to(tl.pointer_type(tl.uint64)) - word_count = bytes_per_descriptor // 8 - for item_index in tl.static_range(0, ITEMS_PER_PROGRAM): - word_offsets = offsets + item_index * BLOCK_SIZE - mask = word_offsets < word_count - values = tl.load(src_ptr + word_offsets, mask=mask, cache_modifier=".cg") - tl.store(dst_ptr + word_offsets, values, mask=mask, cache_modifier=".cs") - - -@torch.no_grad() -def eplb_push_copy(src_ptrs: torch.Tensor, dst_ptrs: torch.Tensor, bytes_per_descriptor: int) -> None: - """Copy 16-byte-aligned expert rows from source to destination pointers.""" - if bytes_per_descriptor <= 64 * 1024: - block_size = 128 - num_warps = 4 - elif bytes_per_descriptor >= 4 * 1024 * 1024 and src_ptrs.numel() > 1: - block_size = 512 - num_warps = 8 - else: - block_size = 256 - num_warps = 4 - items_per_program = 4 - words_per_program = block_size * items_per_program - _eplb_push_copy_kernel[(triton.cdiv(bytes_per_descriptor // 8, words_per_program), src_ptrs.numel())]( - src_ptrs, - dst_ptrs, - bytes_per_descriptor, - BLOCK_SIZE=block_size, - ITEMS_PER_PROGRAM=items_per_program, - num_warps=num_warps, - ) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py index eabe6a6311..8544b22e32 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py @@ -4,7 +4,13 @@ import triton.language as tl from triton.language.standard import _log2, sum, zeros_like -from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_kernels import eplb_replica_index + +@triton.jit +def _eplb_replica_index(token_index, logical_id, replica_count): + """Choose a replica with independent phases for a token's top-k experts.""" + token_hash = token_index.to(tl.uint32) * 2654435769 + expert_hash = logical_id.to(tl.uint32) * 2246822519 + return (token_hash + expert_hash) % replica_count.to(tl.uint32) @triton.jit @@ -344,7 +350,7 @@ def grouped_topk_eplb_kernel( mask=topk_mask, other=1, ) - replica_indices = eplb_replica_index(token_index, selected_logical_ids, replica_counts) + replica_indices = _eplb_replica_index(token_index, selected_logical_ids, replica_counts) selected_physical_ids = tl.load( logical_to_physical_ptr + selected_logical_ids * MAP_SLOTS + replica_indices, mask=topk_mask, diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index aa597e6a11..fbff098acf 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -385,6 +385,7 @@ def _evaluate_after_event(self, event: torch.cuda.Event): ) result["metadata"] = metadata result["layer_plans"] = layer_plans + result["prepared_batches"] = self.transfer.prepare_transfer(layer_plans) with self._evaluation_lock: self._evaluation_result = result except BaseException as exc: @@ -500,7 +501,7 @@ def _start_rebalance(self, result): self.in_flight_layers = [layer_index for layer_index, _ in layer_plans] self.in_flight = True self.in_flight_started_at = time.time() - self.transfer.start(layer_plans) + self.transfer.start(layer_plans, result["prepared_batches"]) if self.global_rank == 0: actual_changed_slot_count = sum(len(plan) for _, plan in layer_plans) cross_node_transfer_count = sum( diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index f0e00df81d..7e5a553571 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -1,18 +1,16 @@ """Asynchronous expert-row migration for EPLB.""" +import ctypes import os +import re import socket import threading from collections import defaultdict, deque from dataclasses import dataclass -from typing import Dict, List, Sequence, Tuple +from typing import Dict, Iterable, List, Optional, Sequence, Tuple import torch import torch.distributed as dist -from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_kernels import ( - eplb_push_copy, -) - @dataclass(frozen=True) class TransferStep: @@ -22,40 +20,6 @@ class TransferStep: src_local_row: int -def extract_expert_tensors(weight) -> List[Tuple[str, torch.Tensor]]: - result = [] - for pack_name in ("w13", "w2"): - pack = getattr(weight, pack_name) - for value_name in ("weight", "weight_scale", "weight_zero_point"): - tensor = getattr(pack, value_name, None) - if tensor is not None: - assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" - result.append((f"{pack_name}.{value_name}", tensor)) - return result - - -def commit_staging_rows( - live: torch.Tensor, - staging: torch.Tensor, - num_experts_per_rank: int, - changed_dst_slots: Sequence[int], -) -> None: - slots = sorted(set(changed_dst_slots)) - if not slots: - return - run_start = previous = slots[0] - for dst_slot in (*slots[1:], None): - if dst_slot is not None and dst_slot == previous + 1: - previous = dst_slot - continue - run_length = previous - run_start + 1 - live.narrow(0, num_experts_per_rank + run_start, run_length).copy_( - staging.narrow(0, run_start, run_length), non_blocking=True - ) - if dst_slot is not None: - run_start = previous = dst_slot - - def align_target_placement(current: torch.Tensor, target: torch.Tensor) -> torch.Tensor: """Canonicalize a target row layout without moving retained experts. @@ -142,7 +106,7 @@ def __init__(self, weights, transfer_group, global_rank, world_size): self.world_size = world_size self.num_experts_per_rank = weights[0].expert_parallel_state.num_primary_experts_per_rank self.device = weights[0].w13.weight.device - self.live = [extract_expert_tensors(weight) for weight in weights] + self.live = [_extract_expert_tensors(weight) for weight in weights] self._validate_live_layout(weights) num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank self.staging = [ @@ -181,12 +145,22 @@ def _validate_live_layout(self, weights) -> None: state.num_redundant_experts_per_rank == num_redundant_slots_per_rank ), "EPLB redundant slot count must match" - def _copy_layer(self, layer_index: int, plan: Sequence[TransferStep], staging) -> None: + def _copy_batch(self, batch, prepared_batch) -> None: raise NotImplementedError - def _copy_batch(self, batch) -> None: - for layer_index, plan, _, staging in batch: - self._copy_layer(layer_index, plan, staging) + def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): + return [ + [ + (layer_index, plan, buffer_index, self.staging[buffer_index]) + for buffer_index, (layer_index, plan) in enumerate( + layer_plans[batch_start : batch_start + self.staging_depth] + ) + ] + for batch_start in range(0, len(layer_plans), self.staging_depth) + ] + + def prepare_transfer(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): + return [(batch, None) for batch in self._make_batches(layer_plans)] def _start_transfer_generation(self) -> None: """Prepare backend state after the in-flight worker check succeeds.""" @@ -194,9 +168,15 @@ def _start_transfer_generation(self) -> None: def _finish_transfer_generation(self) -> None: """Release backend state only after the migration worker has joined.""" - def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]) -> None: + def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepared_batches=None) -> None: if self._thread is not None and self._thread.is_alive(): raise RuntimeError("EPLB transfer is already in flight") + if prepared_batches is None: + prepared_batches = self.prepare_transfer(layer_plans) + expected_batch_count = (len(layer_plans) + self.staging_depth - 1) // self.staging_depth + if len(prepared_batches) != expected_batch_count: + raise ValueError("EPLB prepared batch count does not match layer-plan batches") + # 预构造批次与描述符一起传入,避免推理线程重建。 self._start_transfer_generation() self._error = None with self._pending_lock: @@ -207,11 +187,9 @@ def worker() -> None: torch.cuda.set_device(self.device) if not layer_plans: self._finish_transfer_generation() - for batch_start in range(0, len(layer_plans), self.staging_depth): - batch = [] - for plan_index in range(batch_start, min(batch_start + self.staging_depth, len(layer_plans))): - layer_index, plan = layer_plans[plan_index] - buffer_index = plan_index % self.staging_depth + for batch_index, (batch, prepared_batch) in enumerate(prepared_batches): + batch_start = batch_index * self.staging_depth + for layer_index, plan, buffer_index, _ in batch: release = self._release[buffer_index] # A buffer cannot be reused until its prior committed rows are no longer read by CUDA. release.wait() @@ -221,12 +199,11 @@ def worker() -> None: self._changed_dst_slots[buffer_index] = tuple( step.dst_slot for step in plan if step.dst_rank == self.global_rank ) - batch.append((layer_index, plan, buffer_index, self.staging[buffer_index])) if batch_start > 0 and self._needs_staging_reuse_barrier: # All destinations must finish consuming the prior IPC staging generation # before a source can reuse the peer buffer for this batch. dist.barrier(group=self.transfer_group) - self._copy_batch(batch) + self._copy_batch(batch, prepared_batch) if batch_start + self.staging_depth >= len(layer_plans): self._finish_transfer_generation() with self._pending_lock: @@ -250,7 +227,7 @@ def commit(self, layer_index: int, buffer_index: int, post_copy=None) -> None: self._pending.popleft() changed_dst_slots = self._changed_dst_slots[buffer_index] for (_, live), (_, staging) in zip(self.live[layer_index], self.staging[buffer_index]): - commit_staging_rows( + _commit_staging_rows( live, staging, self.num_experts_per_rank, @@ -277,10 +254,16 @@ class NixlEPLBTransfer(_EPLBTransferBase): """GPU-direct UCX/NIXL EPLB transfer. Initialization errors are fatal.""" backend = "nixl" - staging_depth = 8 _DEFAULT_UCX_TLS = "self,sm,cuda_ipc,cuda_copy,rc_x" + @dataclass + class _PreparedBatch: + remote_entries: Dict[int, list] + push_batch: "_PreparedCudaMemcpyBatch | None" + def __init__(self, weights, transfer_group, global_rank, world_size): + # Reuse at most eight layer buffers to bound EPLB staging memory. + self.staging_depth = min(8, len(weights)) super().__init__(weights, transfer_group, global_rank, world_size) self._nixl_agent = None self._registered_descs = None @@ -292,10 +275,10 @@ def __init__(self, weights, transfer_group, global_rank, world_size): self._same_node_ranks = set() self._cross_node_ranks = set() self._push_stream = torch.cuda.Stream(device=self.device) - self._push_descriptor_cache = {} - self._used_push_descriptor_cache_keys = set() + self._batch_memcpy = _CudaBatchMemcpy() try: self._init_ipc_metadata() + self._init_push_layouts() if self._cross_node_ranks: os.environ.setdefault("UCX_TLS", self._DEFAULT_UCX_TLS) try: @@ -330,11 +313,6 @@ def _init_ipc_metadata(self) -> None: self._needs_staging_reuse_barrier = len(set(hostnames)) < len(hostnames) self._same_node_ranks = {rank for rank, hostname in enumerate(hostnames) if hostname == local_hostname} self._cross_node_ranks = set(range(self.world_size)) - self._same_node_ranks - for layer in self.live: - for name, tensor in layer: - if name.endswith(".weight") and tensor[0].nbytes % 16: - raise RuntimeError(f"NIXL source-push requires 16-byte aligned weight rows: {name}") - from lightllm.server.router.model_infer.mode_backend.pd.p2p_fix import ( p2p_fix_rebuild_cuda_tensor, reduce_tensor, @@ -457,48 +435,58 @@ def _remote_read_cache_key(src_rank: int, entries): ), ) - def _push_staging(self, dst_rank: int, buffer_index: int): - return self.staging[buffer_index] if dst_rank == self.global_rank else self._ipc_staging[dst_rank][buffer_index] + def _init_push_layouts(self) -> None: + self._live_row_layout = [ + [(name, tensor.data_ptr(), tensor[0].nbytes) for name, tensor in layer] for layer in self.live + ] + reference = [(name, row_nbytes) for name, _, row_nbytes in self._live_row_layout[0]] + self._push_staging_row_layout = {} + for dst_rank in self._same_node_ranks: + layouts = [] + for buffer_index in range(self.staging_depth): + staging = ( + self.staging[buffer_index] + if dst_rank == self.global_rank + else self._ipc_staging[dst_rank][buffer_index] + ) + layout = [(name, tensor.data_ptr(), tensor[0].nbytes) for name, tensor in staging] + if [(name, row_nbytes) for name, _, row_nbytes in layout] != reference: + raise RuntimeError("NIXL source-push staging row layout mismatch") + layouts.append(layout) + self._push_staging_row_layout[dst_rank] = layouts - def _cached_descriptor_tensors(self, copies): - key = tuple((source.data_ptr(), destination.data_ptr()) for destination, source in copies) - cached = self._push_descriptor_cache.get(key) - if cached is None: - src_ptrs = torch.tensor([source.data_ptr() for _, source in copies], dtype=torch.int64, device=self.device) - dst_ptrs = torch.tensor( - [destination.data_ptr() for destination, _ in copies], dtype=torch.int64, device=self.device - ) - cached = (src_ptrs, dst_ptrs) - self._push_descriptor_cache[key] = cached - self._used_push_descriptor_cache_keys.add(key) - return cached - - def _push_same_node(self, dst_rank: int, entries) -> None: - staging_by_buffer = {buffer_index: self._push_staging(dst_rank, buffer_index) for _, _, buffer_index in entries} - weight_groups = defaultdict(list) - small_copies = [] - for layer_index, run, buffer_index in entries: - staging = staging_by_buffer[buffer_index] - source_layer = self.live[layer_index] - first = run[0] - run_len = len(run) - for (name, source_tensor), (staging_name, staging_tensor) in zip(source_layer, staging): - if name != staging_name: - raise RuntimeError("NIXL source-push staging tensor name mismatch") - source_rows = source_tensor.narrow(0, first.src_local_row, run_len) - destination_rows = staging_tensor.narrow(0, first.dst_slot, run_len) - if name.endswith(".weight"): - if destination_rows.nbytes % 16: - raise RuntimeError(f"NIXL source-push requires 16-byte aligned weight rows: {name}") - weight_groups[destination_rows.nbytes].append((destination_rows, source_rows)) - else: - small_copies.append((destination_rows, source_rows)) - with torch.cuda.stream(self._push_stream): - for nbytes, copies in weight_groups.items(): - src_ptrs, dst_ptrs = self._cached_descriptor_tensors(copies) - eplb_push_copy(src_ptrs, dst_ptrs, nbytes) - for destination_rows, source_rows in small_copies: - destination_rows.copy_(source_rows, non_blocking=True) + def _prepare_batch(self, batch): + remote_entries = defaultdict(list) + push_descriptors = [] + for layer_index, plan, buffer_index, staging in batch: + steps_by_source = defaultdict(list) + by_destination = defaultdict(list) + for step in plan: + if step.dst_rank == self.global_rank and step.src_rank not in self._same_node_ranks: + steps_by_source[step.src_rank].append(step) + if step.src_rank == self.global_rank and step.dst_rank in self._same_node_ranks: + by_destination[step.dst_rank].append(step) + for src_rank, steps in steps_by_source.items(): + remote_entries[src_rank].extend((layer_index, run, staging) for run in self._contiguous_runs(steps)) + source_layout = self._live_row_layout[layer_index] + for dst_rank, steps in by_destination.items(): + destination_layout = self._push_staging_row_layout[dst_rank][buffer_index] + for run in self._contiguous_runs(steps): + first = run[0] + run_len = len(run) + for (_, source_ptr, row_nbytes), (_, destination_ptr, _) in zip(source_layout, destination_layout): + push_descriptors.append( + ( + source_ptr + first.src_local_row * row_nbytes, + destination_ptr + first.dst_slot * row_nbytes, + run_len * row_nbytes, + ) + ) + push_batch = self._batch_memcpy.prepare(push_descriptors) if push_descriptors else None + return self._PreparedBatch(dict(remote_entries), push_batch) + + def prepare_transfer(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): + return [(batch, self._prepare_batch(batch)) for batch in self._make_batches(layer_plans)] def _get_remote_read(self, src_rank: int, entries): cache_key = self._remote_read_cache_key(src_rank, entries) @@ -557,31 +545,12 @@ def _get_remote_read(self, src_rank: int, entries): self._release_xfers([(local_dlist, remote_dlist, xfer)]) raise - def _copy_batch(self, batch) -> None: - remote_entries = defaultdict(list) - push_entries = defaultdict(list) - for layer_index, plan, _, staging in batch: - steps_by_source = defaultdict(list) - for step in plan: - if step.dst_rank == self.global_rank: - steps_by_source[step.src_rank].append(step) - for src_rank, steps in steps_by_source.items(): - entries = [(layer_index, run, staging) for run in self._contiguous_runs(steps)] - if src_rank not in self._same_node_ranks: - remote_entries[src_rank].extend(entries) - # Source rank owns node-local copies. All ranks build the same batch, - # so buffer_index is the receiver's staging depth index on every peer. - for layer_index, plan, buffer_index, _ in batch: - by_destination = defaultdict(list) - for step in plan: - if step.src_rank == self.global_rank and step.dst_rank in self._same_node_ranks: - by_destination[step.dst_rank].append(step) - for dst_rank, steps in by_destination.items(): - push_entries[dst_rank].extend((layer_index, run, buffer_index) for run in self._contiguous_runs(steps)) - for dst_rank, entries in push_entries.items(): - self._push_same_node(dst_rank, entries) - - xfers = [self._get_remote_read(src_rank, entries) for src_rank, entries in remote_entries.items()] + def _copy_batch(self, batch, prepared_batch) -> None: + if prepared_batch.push_batch is not None: + self._batch_memcpy.enqueue(prepared_batch.push_batch, self._push_stream.cuda_stream) + xfers = [ + self._get_remote_read(src_rank, entries) for src_rank, entries in prepared_batch.remote_entries.items() + ] self._wait_xfers(xfers) self._push_stream.synchronize() # Before a rank publishes this batch it has completed its outgoing source-pushes and @@ -590,7 +559,6 @@ def _copy_batch(self, batch) -> None: def _start_transfer_generation(self) -> None: self._used_xfer_cache_keys.clear() - self._used_push_descriptor_cache_keys.clear() def _finish_transfer_generation(self) -> None: errors = [] @@ -605,8 +573,6 @@ def _finish_transfer_generation(self) -> None: errors.append(exc) else: del self._xfer_cache[cache_key] - for cache_key in set(self._push_descriptor_cache) - self._used_push_descriptor_cache_keys: - del self._push_descriptor_cache[cache_key] if errors: raise RuntimeError("NIXL EPLB cache eviction failed") from errors[0] @@ -614,7 +580,6 @@ def shutdown(self) -> None: agent = self._nixl_agent errors = [] getattr(self, "_used_xfer_cache_keys", set()).clear() - getattr(self, "_used_push_descriptor_cache_keys", set()).clear() if agent is not None: for cache_key, xfer in list(self._xfer_cache.items()): try: @@ -644,7 +609,6 @@ def shutdown(self) -> None: self._registered_descs = None self._nixl_agent = None getattr(self, "_ipc_staging", {}).clear() - getattr(self, "_push_descriptor_cache", {}).clear() if errors: raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] @@ -653,3 +617,169 @@ def __del__(self): self.shutdown() except Exception: pass + + +def _extract_expert_tensors(weight) -> List[Tuple[str, torch.Tensor]]: + result = [] + for pack_name in ("w13", "w2"): + pack = getattr(weight, pack_name) + for value_name in ("weight", "weight_scale", "weight_zero_point"): + tensor = getattr(pack, value_name, None) + if tensor is not None: + assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" + result.append((f"{pack_name}.{value_name}", tensor)) + return result + + +def _commit_staging_rows( + live: torch.Tensor, + staging: torch.Tensor, + num_experts_per_rank: int, + changed_dst_slots: Sequence[int], +) -> None: + slots = sorted(set(changed_dst_slots)) + if not slots: + return + run_start = previous = slots[0] + for dst_slot in (*slots[1:], None): + if dst_slot is not None and dst_slot == previous + 1: + previous = dst_slot + continue + run_length = previous - run_start + 1 + live.narrow(0, num_experts_per_rank + run_start, run_length).copy_( + staging.narrow(0, run_start, run_length), non_blocking=True + ) + if dst_slot is not None: + run_start = previous = dst_slot + + +class _CudaMemLocation(ctypes.Structure): + _fields_ = [("type", ctypes.c_int), ("id", ctypes.c_int)] + + +class _CudaMemcpyAttributes(ctypes.Structure): + _fields_ = [ + ("srcAccessOrder", ctypes.c_int), + ("srcLocHint", _CudaMemLocation), + ("dstLocHint", _CudaMemLocation), + ("flags", ctypes.c_uint), + ] + + +@dataclass +class _PreparedCudaMemcpyBatch: + """Host-side arrays retained for one cudaMemcpyBatchAsync submission.""" + + dsts: object + srcs: object + sizes: object + attrs: _CudaMemcpyAttributes + attrs_idxs: object + count: int + + +class _CudaBatchMemcpy: + """CUDA 13.x ``cudaMemcpyBatchAsync`` binding for EPLB source-push.""" + + _SRC_ACCESS_ORDER_STREAM = 1 + _PREFER_OVERLAP_WITH_COMPUTE = 1 + _CUDA_13_0 = 13000 + _CUDA_14_0 = 14000 + + def __init__(self, library=None): + if library is None: + path = self._find_loaded_cudart() + if path is None: + raise RuntimeError( + "NIXL same-node source-push requires CUDA Runtime 13.x cudaMemcpyBatchAsync; " + "libcudart.so.13 is not loaded" + ) + try: + library = ctypes.CDLL(path) + except OSError as exc: + raise RuntimeError(f"cannot load libcudart: {exc}") from exc + + try: + runtime_get_version = library.cudaRuntimeGetVersion + self._batch_async = library.cudaMemcpyBatchAsync + self._get_error_string = library.cudaGetErrorString + except AttributeError as exc: + raise RuntimeError("cudaMemcpyBatchAsync is unavailable") from exc + + runtime_get_version.restype = ctypes.c_int + runtime_get_version.argtypes = [ctypes.POINTER(ctypes.c_int)] + self._get_error_string.restype = ctypes.c_char_p + self._get_error_string.argtypes = [ctypes.c_int] + runtime_version = ctypes.c_int() + result = runtime_get_version(ctypes.byref(runtime_version)) + if result != 0: + raise RuntimeError(f"cudaRuntimeGetVersion failed with CUDA error {result}") + if not self._CUDA_13_0 <= runtime_version.value < self._CUDA_14_0: + raise RuntimeError( + f"cudaMemcpyBatchAsync requires CUDA Runtime 13.x (13.0 ABI), found {runtime_version.value}" + ) + + pointer_array = ctypes.POINTER(ctypes.c_void_p) + self._batch_async.restype = ctypes.c_int + self._batch_async.argtypes = [ + pointer_array, + pointer_array, + ctypes.POINTER(ctypes.c_size_t), + ctypes.c_size_t, + ctypes.POINTER(_CudaMemcpyAttributes), + ctypes.POINTER(ctypes.c_size_t), + ctypes.c_size_t, + ctypes.c_void_p, + ] + + @staticmethod + def prepare(copies: Iterable[Tuple[int, int, int]]) -> _PreparedCudaMemcpyBatch: + copies = tuple(copies) + if not copies: + raise ValueError("cudaMemcpyBatchAsync requires at least one copy") + for src, dst, size in copies: + if not src or not dst or size <= 0: + raise ValueError("cudaMemcpyBatchAsync requires non-null pointers and positive sizes") + count = len(copies) + dsts = (ctypes.c_void_p * count)(*(dst for _, dst, _ in copies)) + srcs = (ctypes.c_void_p * count)(*(src for src, _, _ in copies)) + sizes = (ctypes.c_size_t * count)(*(size for _, _, size in copies)) + attrs = _CudaMemcpyAttributes() + attrs.srcAccessOrder = _CudaBatchMemcpy._SRC_ACCESS_ORDER_STREAM + attrs.flags = _CudaBatchMemcpy._PREFER_OVERLAP_WITH_COMPUTE + attrs_idxs = (ctypes.c_size_t * 1)(0) + return _PreparedCudaMemcpyBatch(dsts, srcs, sizes, attrs, attrs_idxs, count) + + def enqueue(self, prepared: _PreparedCudaMemcpyBatch, stream: int) -> None: + result = self._batch_async( + prepared.dsts, + prepared.srcs, + prepared.sizes, + prepared.count, + ctypes.byref(prepared.attrs), + prepared.attrs_idxs, + 1, + ctypes.c_void_p(stream), + ) + if result != 0: + message = self._get_error_string(result) + error = message.decode("utf-8") if message else f"CUDA error {result}" + raise RuntimeError(f"cudaMemcpyBatchAsync failed: {error}") + + @staticmethod + def _find_loaded_cudart() -> Optional[str]: + """Return a mapped CUDA 13 runtime without loading CUDA as a side effect.""" + try: + with open("/proc/self/maps") as maps: + for line in maps: + if "libcudart" not in line: + continue + path_start = line.find("/") + if path_start < 0: + continue + path = line[path_start:].strip().removesuffix(" (deleted)") + if re.search(r"libcudart[^/]*\.so\.13(?:\D|$)", path): + return path + except OSError: + pass + return None diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 6105cd0724..e34995a53f 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1,3 +1,5 @@ +import builtins +import io import threading import time from collections import deque @@ -48,10 +50,11 @@ ) from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( TransferStep, + _CudaBatchMemcpy, + _commit_staging_rows, + _extract_expert_tensors, align_target_placement, build_transfer_plan, - commit_staging_rows, - extract_expert_tensors, ) @@ -1536,6 +1539,14 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c manager._evaluation_result = None manager._evaluation_error = None manager._continuous_collection_end_step = None + prepare_calls = [] + prepared_batches = object() + + def prepare_transfer(layer_plans): + prepare_calls.append(layer_plans) + return prepared_batches + + manager.transfer = SimpleNamespace(prepare_transfer=prepare_transfer) manager._collect_local_samples = lambda: torch.full((1, 3, 4), 100, dtype=torch.int64) planned_placement = torch.tensor( [ @@ -1572,6 +1583,8 @@ def build_maps_for_layers(*args, **kwargs): metadata = manager._evaluation_result["metadata"] assert metadata[1] is None assert [layer_index for layer_index, _plan in manager._evaluation_result["layer_plans"]] == [0, 2] + assert prepare_calls == [manager._evaluation_result["layer_plans"]] + assert manager._evaluation_result["prepared_batches"] is prepared_batches for layer_index in (0, 2): item = metadata[layer_index] expected = build_logical_to_physical_map( @@ -1584,6 +1597,57 @@ def build_maps_for_layers(*args, **kwargs): assert torch.equal(item[1], expected[1]) +def test_manager_preparation_error_is_saved_as_evaluation_error(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=torch.zeros((1, 2), dtype=torch.int64), + num_logical_experts=2, + world_size=1, + ) + }, + )() + ] + manager._eplb_states = [manager.weights[0].expert_parallel_state.eplb] + manager.global_rank = 0 + manager.world_size = 1 + manager.node_world_size = 1 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.num_logical_experts = 2 + manager.current_placement = torch.tensor([[[0], [1]]], dtype=torch.int64) + manager.evaluation_group = object() + manager._evaluation_lock = threading.Lock() + manager._evaluation_result = None + manager._evaluation_error = None + manager._continuous_collection_end_step = None + manager._collect_local_samples = lambda: torch.ones((1, 1, 2), dtype=torch.int64) + manager._plan_and_broadcast = lambda _global_load: { + "kind": "planned", + "placement": torch.tensor([[[0], [1]]], dtype=torch.int64), + "improved": torch.tensor([True]), + } + + def fail_prepare(_layer_plans): + raise RuntimeError("prepare failed") + + manager.transfer = SimpleNamespace(prepare_transfer=fail_prepare) + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) + monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) + monkeypatch.setattr(manager_module, "build_transfer_plan", lambda *_args, **_kwargs: []) + + manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) + + assert manager._evaluation_result is None + assert isinstance(manager._evaluation_error, RuntimeError) + assert str(manager._evaluation_error) == "prepare failed" + + def test_decode_dispatch_keeps_logical_ids_and_uses_logical_expert_count(monkeypatch): class Buffer: def low_latency_dispatch(self, **kwargs): @@ -1941,7 +2005,7 @@ def __init__(self, offset, scale=True, zero=True): self.weight_zero_point = torch.full((3, 1), offset + 2) if zero else None weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False, zero=False)})() - tensors = extract_expert_tensors(weight) + tensors = _extract_expert_tensors(weight) assert [name for name, _ in tensors] == [ "w13.weight", "w13.weight_scale", @@ -1953,7 +2017,7 @@ def __init__(self, offset, scale=True, zero=True): def test_commit_staging_rows_only_overwrites_redundant_rows(): live = torch.arange(20).reshape(5, 4) staging = torch.full((2, 4), -1) - commit_staging_rows( + _commit_staging_rows( live, staging, num_experts_per_rank=3, @@ -1968,7 +2032,7 @@ def test_commit_staging_rows_preserves_unchanged_destination_slots(): staging = torch.tensor([[-1, -1, -1, -1], [-2, -2, -2, -2], [-3, -3, -3, -3], [-4, -4, -4, -4]]) original = live.clone() - commit_staging_rows( + _commit_staging_rows( live, staging, num_experts_per_rank=3, @@ -2010,7 +2074,7 @@ def __init__(self, owner, rows): def narrow(self, _dim, start, length): return View(self.owner, start, length) - commit_staging_rows( + _commit_staging_rows( Tensor("live", 20), Tensor("staging", 4), num_experts_per_rank=10, @@ -2291,11 +2355,12 @@ def synchronize(self): transfer.transfer_group = "transfer-group" copied = [] - def copy_layer(layer, _plan, _staging): - copied.append(layer) - operations.append(("copy", layer)) + def copy_batch(batch, _prepared_batch): + for layer, _plan, _buffer, _staging in batch: + copied.append(layer) + operations.append(("copy", layer)) - transfer._copy_layer = copy_layer + transfer._copy_batch = copy_batch monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda device: None) monkeypatch.setattr(transfer_module.torch.cuda, "Event", Event) monkeypatch.setattr(transfer_module.torch.cuda, "current_stream", lambda: object()) @@ -2305,7 +2370,10 @@ def copy_layer(layer, _plan, _staging): lambda **kwargs: operations.append(("barrier", kwargs["group"])), ) - transfer.start([(0, []), (1, []), (2, [])]) + plans = [(0, []), (1, []), (2, [])] + prepared_batches = transfer.prepare_transfer(plans) + monkeypatch.setattr(transfer, "_make_batches", lambda _plans: pytest.fail("start must reuse prepared batches")) + transfer.start(plans, prepared_batches) deadline = time.monotonic() + 2 while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: time.sleep(0.001) @@ -2349,7 +2417,7 @@ def make_transfer(finalize): transfer._error = None transfer._thread = None transfer._needs_staging_reuse_barrier = False - transfer._copy_batch = lambda _batch: None + transfer._copy_batch = lambda _batch, _prepared_batch: None transfer._finish_transfer_generation = finalize return transfer @@ -2672,15 +2740,21 @@ def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): manager.sampling_interval = 320 manager._continuous_collection_start_step = 0 manager.global_rank = 1 - manager.transfer = type("Transfer", (), {"start": lambda self, plans: setattr(self, "plans", plans)})() + manager.transfer = type( + "Transfer", + (), + {"start": lambda self, plans, prepared_batches: setattr(self, "started", (plans, prepared_batches))}, + )() manager._reset_recorded_samples = lambda: None + prepared_batches = [object()] manager._start_rebalance( { "placement": torch.zeros((1, 1, 1), dtype=torch.int64), "improved": torch.tensor([True]), "metadata": [None], "layer_plans": [(0, object())], + "prepared_batches": prepared_batches, "before": {"max": 1.0, "p95": 1.0}, "after": {"max": 1.0, "p95": 1.0}, "model_imbalance_ratio": 1.0, @@ -2693,6 +2767,18 @@ def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): assert manager.sampling_interval == 20 assert manager.in_flight assert manager._continuous_collection_start_step is None + assert len(manager.transfer.started[0]) == 1 + assert manager.transfer.started[1] is prepared_batches + + +def test_transfer_start_rejects_prepared_batches_with_wrong_batch_count(): + transfer = object.__new__(transfer_module._EPLBTransferBase) + transfer._thread = None + transfer.staging_depth = 2 + transfer.staging = [object(), object()] + + with pytest.raises(ValueError, match="prepared batch count"): + transfer.start([(0, []), (1, []), (2, [])], prepared_batches=[object()]) def test_first_rebalance_completion_switches_to_four_step_sparse_window(monkeypatch): @@ -2780,6 +2866,245 @@ def test_nixl_descriptor_runs_merge_only_jointly_contiguous_source_and_destinati ] +def test_nixl_prepare_batch_compiles_hot_path_without_tensor_views(monkeypatch): + class BatchMemcpy: + def __init__(self): + self.prepared = [] + self.enqueued = [] + + def prepare(self, descriptors): + descriptor = tuple(descriptors) + self.prepared.append(descriptor) + return descriptor + + def enqueue(self, descriptor, stream): + self.enqueued.append((descriptor, stream)) + + stream = SimpleNamespace(cuda_stream=123, synchronize=lambda: None) + batch_memcpy = BatchMemcpy() + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer._push_stream = stream + transfer._batch_memcpy = batch_memcpy + transfer.global_rank = 0 + transfer._same_node_ranks = {0, 1, 2} + transfer._live_row_layout = [ + [ + ("w13.weight", 1000, 32), + ("w13.weight_scale", 2000, 32), + ("w13.weight_zero_point", 3000, 32), + ] + ] + transfer._push_staging_row_layout = { + 1: [[("w13.weight", 4000, 32), ("w13.weight_scale", 5000, 32), ("w13.weight_zero_point", 6000, 32)]], + 2: [[("w13.weight", 7000, 32), ("w13.weight_scale", 8000, 32), ("w13.weight_zero_point", 9000, 32)]], + } + transfer._get_remote_read = lambda *_args: None + transfer._wait_xfers = lambda _xfers: None + + run_a = [TransferStep(1, 3, 0, 5), TransferStep(1, 4, 0, 6)] + run_b = [TransferStep(2, 1, 0, 2)] + local_inbound = TransferStep(0, 0, 1, 0) + remote_inbound = TransferStep(0, 1, 3, 2) + staging = object() + batch = [(0, run_a + run_b + [local_inbound, remote_inbound], 0, staging)] + prepared = transfer._prepare_batch(batch) + monkeypatch.setattr(transfer, "_prepare_batch", lambda _batch: pytest.fail("hot path must not prepare descriptors")) + monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) + transfer._copy_batch(batch, prepared) + + expected = ( + (1160, 4096, 64), + (2160, 5096, 64), + (3160, 6096, 64), + (1064, 7032, 32), + (2064, 8032, 32), + (3064, 9032, 32), + ) + assert batch_memcpy.prepared == [expected] + assert batch_memcpy.enqueued == [(expected, 123)] + assert prepared.remote_entries == {3: [(0, [remote_inbound], staging)]} + + +def test_nixl_prepare_transfer_batches_match_staging_depth(): + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer.staging_depth = 2 + transfer.staging = ["staging-0", "staging-1"] + seen_batches = [] + + def prepare_batch(batch): + seen_batches.append(batch) + return f"prepared-{len(seen_batches)}" + + transfer._prepare_batch = prepare_batch + layer_plans = [(3, "plan-3"), (4, "plan-4"), (5, "plan-5")] + + prepared_batches = transfer.prepare_transfer(layer_plans) + + assert prepared_batches == [ + (seen_batches[0], "prepared-1"), + (seen_batches[1], "prepared-2"), + ] + assert [[(layer, buffer) for layer, _plan, buffer, _staging in batch] for batch in seen_batches] == [ + [(3, 0), (4, 1)], + [(5, 0)], + ] + + +def test_cuda_batch_memcpy_cuda13_abi_and_descriptor_layout(): + class Function: + def __init__(self, callback): + self.callback = callback + self.restype = None + self.argtypes = None + + def __call__(self, *args): + return self.callback(*args) + + class Library: + def __init__(self): + def get_version(pointer): + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 13000 + return 0 + + self.cudaRuntimeGetVersion = Function(get_version) + self.cudaMemcpyBatchAsync = Function(lambda *args: self.calls.append(args) or 0) + self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + self.calls = [] + + import ctypes + + library = Library() + batch_memcpy = _CudaBatchMemcpy(library) + prepared = batch_memcpy.prepare(((101, 201, 64), (102, 202, 128))) + batch_memcpy.enqueue(prepared, 777) + + assert len(library.cudaMemcpyBatchAsync.argtypes) == 8 + assert library.calls[0][3] == 2 + assert library.calls[0][6] == 1 + assert library.calls[0][7].value == 777 + assert [pointer for pointer in library.calls[0][0]] == [201, 202] + assert [pointer for pointer in library.calls[0][1]] == [101, 102] + assert list(library.calls[0][2]) == [64, 128] + attrs = library.calls[0][4]._obj + assert attrs.srcAccessOrder == 1 + assert attrs.srcLocHint.type == attrs.srcLocHint.id == 0 + assert attrs.dstLocHint.type == attrs.dstLocHint.id == 0 + assert attrs.flags == 1 + + +def test_cuda_batch_memcpy_rejects_unsupported_runtime_and_invalid_descriptors(): + class Function: + def __init__(self, callback): + self.callback = callback + self.restype = None + self.argtypes = None + + def __call__(self, *args): + return self.callback(*args) + + class OldRuntimeLibrary: + def __init__(self): + def get_version(pointer): + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 12080 + return 0 + + self.cudaRuntimeGetVersion = Function(get_version) + self.cudaMemcpyBatchAsync = Function(lambda *_args: 0) + self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + + import ctypes + + with pytest.raises(RuntimeError, match="13.0"): + _CudaBatchMemcpy(OldRuntimeLibrary()) + + class FutureRuntimeLibrary: + def __init__(self): + def get_version(pointer): + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 14000 + return 0 + + self.cudaRuntimeGetVersion = Function(get_version) + self.cudaMemcpyBatchAsync = Function(lambda *_args: 0) + self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + + with pytest.raises(RuntimeError, match="13.x"): + _CudaBatchMemcpy(FutureRuntimeLibrary()) + + class MissingBatchSymbolLibrary: + def __init__(self): + def get_version(pointer): + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 13000 + return 0 + + self.cudaRuntimeGetVersion = Function(get_version) + self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + + with pytest.raises(RuntimeError, match="cudaMemcpyBatchAsync"): + _CudaBatchMemcpy(MissingBatchSymbolLibrary()) + with pytest.raises(ValueError, match="at least one"): + _CudaBatchMemcpy.prepare(()) + with pytest.raises(ValueError, match="positive"): + _CudaBatchMemcpy.prepare(((1, 2, 0),)) + + +def test_nixl_transfer_fails_fast_without_cuda13_batch_memcpy(monkeypatch): + failure = RuntimeError("missing cudaMemcpyBatchAsync") + + def unavailable(): + raise failure + + monkeypatch.setattr( + transfer_module._EPLBTransferBase, "__init__", lambda self, *_args: setattr(self, "device", "mock") + ) + monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) + monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", unavailable) + with pytest.raises(RuntimeError, match="missing cudaMemcpyBatchAsync") as exc_info: + transfer_module.NixlEPLBTransfer([object()], object(), 0, 1) + assert exc_info.value is failure + + +@pytest.mark.parametrize( + ("maps", "open_raises", "expected"), + [ + ( + "7f /cuda-12/libcudart.so.12.8\n" + "7f /first/libcudart.so.13 (deleted)\n" + "7f /second/libcudart.so.13\n" + "7f /cuda-14/libcudart.so.14\n" + "7f /cuda-130/libcudart.so.130\n", + False, + "/first/libcudart.so.13", + ), + ("7f /cuda-12/libcudart.so.12\n7f /cuda-130/libcudart.so.130\n", False, None), + ("", True, None), + ], +) +def test_cuda_batch_memcpy_finds_first_loaded_cuda13_runtime(monkeypatch, maps, open_raises, expected): + def open_maps(_path): + if open_raises: + raise OSError("maps unavailable") + return io.StringIO(maps) + + monkeypatch.setattr(builtins, "open", open_maps) + assert _CudaBatchMemcpy._find_loaded_cudart() == expected + + +@pytest.mark.parametrize(("layers", "expected_depth"), [(1, 1), (8, 8), (9, 8), (43, 8)]) +def test_nixl_transfer_bounds_staging_depth(monkeypatch, layers, expected_depth): + def base_init(self, *_args): + self.device = "mock-device" + + monkeypatch.setattr(transfer_module._EPLBTransferBase, "__init__", base_init) + monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) + monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", lambda: object()) + monkeypatch.setattr(transfer_module.NixlEPLBTransfer, "_init_ipc_metadata", lambda self: None) + monkeypatch.setattr(transfer_module.NixlEPLBTransfer, "_init_push_layouts", lambda self: None) + + transfer = transfer_module.NixlEPLBTransfer([object()] * layers, object(), 0, 1) + + assert transfer.staging_depth == expected_depth + + def test_nixl_remote_read_cache_reuses_exact_batch_key_and_releases_on_shutdown(): class Tensor: nbytes = 8 @@ -2853,7 +3178,7 @@ def remove_remote_agent(self, remote_name): assert agent.removed_agents == ["remote-1"] -def test_nixl_descriptor_caches_are_bounded_to_the_current_transfer_generation( +def test_nixl_remote_read_cache_is_bounded_to_the_current_transfer_generation( monkeypatch, ): class Tensor: @@ -2904,8 +3229,6 @@ def remove_remote_agent(self, _): transfer._nixl_agent = agent transfer._xfer_cache = {} transfer._used_xfer_cache_keys = set() - transfer._push_descriptor_cache = {} - transfer._used_push_descriptor_cache_keys = set() transfer._remote_agents = {1: "remote-1"} transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} transfer._registered_descs = None @@ -2932,53 +3255,38 @@ def remove_remote_agent(self, _): staging = [("weight", Tensor(200))] entries_a = [(0, [TransferStep(0, 0, 1, 0)], staging)] entries_b = [(0, [TransferStep(0, 1, 1, 0)], staging)] - source_a, destination_a = Tensor(300), Tensor(400) - source_b, destination_b = Tensor(500), Tensor(600) - copies_a = [(destination_a, source_a)] - copies_b = [(destination_b, source_b)] - key_a = ((source_a.data_ptr(), destination_a.data_ptr()),) - key_b = ((source_b.data_ptr(), destination_b.data_ptr()),) - stale_key = ((700, 800),) - transfer._push_descriptor_cache = { - key_a: (object(), object()), - stale_key: (object(), object()), - } - generation = [entries_a, copies_a] + generation = [entries_a] - def copy_batch(_batch): + def copy_batch(_batch, _prepared_batch): transfer._get_remote_read(1, generation[0]) - transfer._cached_descriptor_tensors(generation[1]) transfer._copy_batch = copy_batch monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) - transfer.start([(0, [])]) + prepared_batches = [([(0, [], 0, transfer.staging[0])], None)] + transfer.start([(0, [])], prepared_batches) transfer.finish() assert agent.made == 1 assert agent.released_xfers == agent.released_dlists == 0 - assert set(transfer._push_descriptor_cache) == {key_a} assert len(transfer._xfer_cache) == 1 # The real manager releases each staging buffer through commit(). This # focused cache test has no commits, so model that hand-off before the # next generation reuses buffer zero. transfer._release[0].set() - transfer.start([(0, [])]) + transfer.start([(0, [])], prepared_batches) transfer.finish() assert agent.made == 1 assert agent.released_xfers == agent.released_dlists == 0 - assert set(transfer._push_descriptor_cache) == {key_a} assert len(transfer._xfer_cache) == 1 - transfer._push_descriptor_cache[key_b] = (object(), object()) - generation[:] = [entries_b, copies_b] + generation[:] = [entries_b] transfer._release[0].set() - transfer.start([(0, [])]) + transfer.start([(0, [])], prepared_batches) transfer.finish() assert agent.made == 2 assert agent.released_xfers == 1 assert agent.released_dlists == 2 - assert set(transfer._push_descriptor_cache) == {key_b} assert len(transfer._xfer_cache) == 1 transfer.shutdown() @@ -3051,22 +3359,16 @@ def test_nixl_copy_batch_source_pushes_local_rows_and_keeps_remote_ucx_reads( class Stream: def __init__(self): self.synchronized = 0 + self.cuda_stream = 123 def synchronize(self): self.synchronized += 1 - @contextmanager - def use_stream(_stream): - yield - stream = Stream() transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer.global_rank = 0 - transfer.num_experts_per_rank = 2 - transfer._same_node_ranks = {0, 1} transfer._push_stream = stream - remote_reads, pushed, waited_xfers = [], [], [] - transfer._push_same_node = lambda dst_rank, entries: pushed.append((dst_rank, entries)) + transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) + remote_reads, waited_xfers, enqueued = [], [], [] transfer._get_remote_read = lambda src_rank, entries: remote_reads.append((src_rank, entries)) or ( None, None, @@ -3074,16 +3376,14 @@ def use_stream(_stream): ) transfer._wait_xfers = waited_xfers.extend - monkeypatch.setattr(transfer_module.torch.cuda, "stream", use_stream) - staging = [] - local_step = TransferStep(1, 0, 0, 0) - remote_step = TransferStep(0, 1, 2, 0) - - transfer._copy_batch([(0, [local_step], 0, staging)]) - assert pushed == [(1, [(0, [local_step], 0)])] - assert remote_reads == [] + monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) + prepared_push = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") + transfer._copy_batch([], prepared_push) + assert enqueued == [("push", stream.cuda_stream)] - transfer._copy_batch([(0, [remote_step], 0, staging)]) + remote_step = TransferStep(0, 1, 2, 0) + prepared_remote = transfer_module.NixlEPLBTransfer._PreparedBatch({2: [(0, [remote_step], [])]}, None) + transfer._copy_batch([], prepared_remote) assert [rank for rank, _ in remote_reads] == [2] assert waited_xfers == [(None, None, "xfer")] @@ -3092,28 +3392,22 @@ def test_nixl_copy_batch_self_only_rank_pushes_and_synchronizes(monkeypatch): class Stream: def __init__(self): self.synchronized = 0 + self.cuda_stream = 456 def synchronize(self): self.synchronized += 1 - @contextmanager - def use_stream(_stream): - yield - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer.global_rank = 0 - transfer.num_experts_per_rank = 2 - transfer._same_node_ranks = {0} transfer._push_stream = Stream() - pushed = [] - transfer._push_same_node = lambda dst_rank, entries: pushed.append((dst_rank, entries)) + enqueued = [] + transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) transfer._wait_xfers = lambda _xfers: None - monkeypatch.setattr(transfer_module.torch.cuda, "stream", use_stream) + monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) - self_step = TransferStep(0, 0, 0, 0) - transfer._copy_batch([(0, [self_step], 0, [])]) + prepared = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") + transfer._copy_batch([], prepared) - assert pushed == [(0, [(0, [self_step], 0)])] + assert enqueued == [("push", transfer._push_stream.cuda_stream)] assert transfer._push_stream.synchronized == 1 diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 4397943be4..46c218ad45 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -1,4 +1,4 @@ -"""Two-GPU NIXL EPLB correctness and 512 MiB micro-performance test.""" +"""NIXL EPLB correctness tests and a two-GPU 512 MiB micro-performance test.""" import os import random import socket @@ -12,6 +12,7 @@ from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( NixlEPLBTransfer, + align_target_placement, build_transfer_plan, ) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( @@ -134,13 +135,15 @@ def _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload): (layer_index, plan, buffer_index, transfer.staging[buffer_index]) for buffer_index, (layer_index, plan) in enumerate(layer_plans) ] + # Measure the precompiled hot path only; planning and descriptor construction are off the timer. + prepared_batch = transfer._prepare_batch(batch) for _ in range(3): - transfer._copy_batch(batch) + transfer._copy_batch(batch, prepared_batch) dist.barrier(group=control_group) samples = [] for _ in range(20): started = time.perf_counter() - transfer._copy_batch(batch) + transfer._copy_batch(batch, prepared_batch) samples.append(payload / (time.perf_counter() - started) / 1e9) torch.cuda.synchronize() median = statistics.median(samples) @@ -223,3 +226,219 @@ def test_eplb_transfer_two_gpu_correctness_and_microperf(): f"NIXL={nixl_gbps:.2f} GB/s NIXL _copy_batch={nixl_copy_batch_gbps:.2f} GB/s" ) assert nixl_gbps > 0 + + +def _depth_value(expert, layer_index, offset): + return expert // 32 * 100 + layer_index * 5 + expert % 32 + offset + + +class _DepthWeight: + def __init__(self, rank, layer_index, initial_placement): + self.n_routed_experts = 256 + self.expert_parallel_state = ExpertParallelState( + num_logical_experts=256, + world_size=8, + eplb=EPLBState( + num_redundant_experts_per_rank=4, + initial_redundant_expert_ids_by_rank=initial_placement.clone(), + logical_to_physical_map=torch.zeros((256, 8), dtype=torch.int32, device="cuda"), + logical_replica_count=torch.ones(256, dtype=torch.int32, device="cuda"), + route_counter=torch.zeros((1, 256), dtype=torch.int64, device="cuda"), + ), + ) + logical_ids = list(range(rank * 32, (rank + 1) * 32)) + initial_placement[rank].tolist() + self.w13 = self._pack(logical_ids, layer_index, 0) + self.w2 = self._pack(logical_ids, layer_index, 2) + + @staticmethod + def _pack(logical_ids, layer_index, offset): + weight = torch.empty((36, 64), dtype=torch.float16, device="cuda") + scale = torch.empty((36, 1), dtype=torch.float32, device="cuda") + for row, expert in enumerate(logical_ids): + value = _depth_value(expert, layer_index, offset) + weight[row].fill_(value) + scale[row].fill_(value + 0.25) + return _Pack(weight, scale) + + +def _depth_target(layer_index): + return torch.tensor([[((dst + layer_index + slot + 1) % 8) * 32 + slot for slot in range(4)] for dst in range(8)]) + + +def _wait_all_pending(transfer, group, expected_count): + deadline = time.monotonic() + 30 + while True: + pending = transfer.pending_layers() + ready_count = torch.tensor([len(pending)], dtype=torch.int32) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=group) + if int(ready_count.item()) == expected_count: + return pending + if time.monotonic() > deadline: + raise TimeoutError(f"expected {expected_count} pending layers, got {pending}") + time.sleep(0.001) + + +def _clone_depth_live(weights): + return [ + [tensor.detach().clone() for _, tensor in transfer_tensors] + for transfer_tensors in [ + [ + ("w13.weight", weight.w13.weight), + ("w13.scale", weight.w13.weight_scale), + ("w2.weight", weight.w2.weight), + ("w2.scale", weight.w2.weight_scale), + ] + for weight in weights + ] + ] + + +def _assert_depth_snapshot(weights, snapshot, layer_indices=None, primary_only=False): + if layer_indices is None: + layer_indices = range(len(weights)) + for layer_index in layer_indices: + live_tensors = ( + weights[layer_index].w13.weight, + weights[layer_index].w13.weight_scale, + weights[layer_index].w2.weight, + weights[layer_index].w2.weight_scale, + ) + for live, expected in zip(live_tensors, snapshot[layer_index]): + if primary_only: + live = live[:32] + expected = expected[:32] + torch.testing.assert_close(live, expected) + + +def _assert_depth_staging(rank, layer_plans, target_placements, pending, transfer): + for buffer_index, ((layer_index, plan), pending_item) in enumerate(zip(layer_plans, pending)): + assert pending_item == (layer_index, buffer_index) + expected = {step.dst_slot: step for step in plan if step.dst_rank == rank} + for dst_slot, step in expected.items(): + expert = int(target_placements[layer_index][rank, dst_slot]) + base = _depth_value(expert, layer_index, 0) + staging = transfer.staging[buffer_index] + assert torch.all(staging[0][1][dst_slot] == base) + assert torch.all(staging[1][1][dst_slot] == base + 0.25) + assert torch.all(staging[2][1][dst_slot] == base + 2) + assert torch.all(staging[3][1][dst_slot] == base + 2.25) + + +def _assert_depth_live(weights, rank, layer_plans, target_placements): + for layer_index, plan in layer_plans: + for step in plan: + if step.dst_rank != rank: + continue + expert = int(target_placements[layer_index][rank, step.dst_slot]) + base = _depth_value(expert, layer_index, 0) + assert torch.all(weights[layer_index].w13.weight[32 + step.dst_slot] == base) + assert torch.all(weights[layer_index].w13.weight_scale[32 + step.dst_slot] == base + 0.25) + assert torch.all(weights[layer_index].w2.weight[32 + step.dst_slot] == base + 2) + assert torch.all(weights[layer_index].w2.weight_scale[32 + step.dst_slot] == base + 2.25) + + +def _assert_peer_coverage(layer_plans, require_redundant_source=False): + steps = [step for _, plan in layer_plans for step in plan] + assert {step.dst_rank for step in steps} == set(range(8)) + assert {step.src_rank for step in steps} == set(range(8)) + for source_rank in range(8): + assert len({step.dst_rank for step in steps if step.src_rank == source_rank}) > 1 + if require_redundant_source: + assert any(step.src_local_row >= 32 for step in steps) + + +def _depth_worker(rank, port): + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = str(port) + torch.cuda.set_device(rank) + dist.init_process_group("gloo", rank=rank, world_size=8) + control_group = dist.new_group(list(range(8)), backend="gloo") + transfer_group = dist.new_group(list(range(8)), backend="gloo") + initial_placement = build_initial_redundant_expert_ids(256, 8, 4) + weights = [_DepthWeight(rank, layer_index, initial_placement) for layer_index in range(9)] + transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=8) + assert transfer.staging_depth == 8 + staging_bytes = sum(tensor.nbytes for staging in transfer.staging for _, tensor in staging) + one_layer_staging_bytes = sum(tensor[:4].nbytes for _, tensor in transfer.live[0]) + assert staging_bytes == 8 * one_layer_staging_bytes + current = initial_placement + first_order = [8, 0, 5, 1, 7, 3, 6, 2, 4] + first_targets = { + layer_index: align_target_placement(current, _depth_target(layer_index)) for layer_index in range(9) + } + first_plans = [ + (layer_index, build_transfer_plan(current, first_targets[layer_index], 256, 8, 8)) + for layer_index in first_order + ] + _assert_peer_coverage(first_plans) + first_snapshot = _clone_depth_live(weights) + transfer.start(first_plans, transfer.prepare_transfer(first_plans)) + committed = 0 + pending = _wait_all_pending(transfer, control_group, 8) + _assert_depth_staging(rank, first_plans[:8], first_targets, pending, transfer) + _assert_depth_snapshot(weights, first_snapshot) + dist.barrier(group=control_group) + if rank == 0: + time.sleep(0.1) + for layer_index, buffer_index in pending: + _assert_depth_snapshot(weights, first_snapshot, [layer for layer, _ in first_plans[committed:]]) + transfer.commit(layer_index, buffer_index) + committed += 1 + pending = _wait_all_pending(transfer, control_group, 1) + _assert_depth_staging(rank, first_plans[8:], first_targets, pending, transfer) + _assert_depth_snapshot(weights, first_snapshot, [first_plans[8][0]]) + dist.barrier(group=control_group) + if rank == 0: + time.sleep(0.1) + for layer_index, buffer_index in pending: + _assert_depth_snapshot(weights, first_snapshot, [layer for layer, _ in first_plans[committed:]]) + transfer.commit(layer_index, buffer_index) + committed += 1 + transfer.finish() + torch.cuda.synchronize() + dist.barrier(group=control_group) + _assert_depth_live(weights, rank, first_plans, first_targets) + + second_order = [7, 2, 4] + second_targets = { + layer_index: align_target_placement( + first_targets[layer_index], + torch.tensor([[((dst + layer_index + slot + 3) % 8) * 32 + slot for slot in range(4)] for dst in range(8)]), + ) + for layer_index in second_order + } + second_plans = [ + ( + layer_index, + build_transfer_plan(first_targets[layer_index], second_targets[layer_index], 256, 8, 8), + ) + for layer_index in second_order + ] + _assert_peer_coverage(second_plans, require_redundant_source=True) + second_snapshot = _clone_depth_live(weights) + transfer.start(second_plans, transfer.prepare_transfer(second_plans)) + pending = _wait_all_pending(transfer, control_group, len(second_plans)) + _assert_depth_staging(rank, second_plans, second_targets, pending, transfer) + _assert_depth_snapshot(weights, second_snapshot) + dist.barrier(group=control_group) + if rank == 0: + time.sleep(0.1) + for layer_index, buffer_index in pending: + transfer.commit(layer_index, buffer_index) + transfer.finish() + torch.cuda.synchronize() + dist.barrier(group=control_group) + _assert_depth_live(weights, rank, second_plans, second_targets) + _assert_depth_snapshot(weights, second_snapshot, set(range(9)) - set(second_order)) + _assert_depth_snapshot(weights, second_snapshot, primary_only=True) + transfer.shutdown() + dist.barrier(group=control_group) + dist.destroy_process_group() + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < 8, + reason="requires eight CUDA GPUs", +) +def test_eplb_transfer_eight_gpu_bounded_staging_reuse(): + mp.spawn(_depth_worker, args=(_free_port(),), nprocs=8, join=True) From 5a7440ad6dfc88632a8e160e0ba834a0db305b4d Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Tue, 8 Sep 2026 14:13:45 +0800 Subject: [PATCH 06/72] refine --- .../layer_infer/cache_tensor_manager.py | 7 +- .../meta_weights/fused_moe/eplb_placement.py | 224 +++++------------- .../meta_weights/fused_moe/impl/__init__.py | 25 +- .../fused_moe/impl/deepgemm_impl.py | 15 +- .../triton_kernel/fused_moe/grouped_topk.py | 1 - lightllm/common/eplb_utils.py | 16 ++ .../kv_cache_mem_manager/mem_manager.py | 43 +++- .../model_infer/mode_backend/base_backend.py | 5 +- .../model_infer/mode_backend/eplb_manager.py | 39 +-- .../model_infer/mode_backend/eplb_transfer.py | 40 +--- lightllm/utils/envs_utils.py | 1 + lightllm/utils/profile_max_tokens.py | 8 + unit_tests/common/fused_moe/test_eplb.py | 163 +++++++------ .../fused_moe/test_eplb_transfer_gpu.py | 2 +- .../models/deepseek_v4/test_memory_profile.py | 64 +++++ 15 files changed, 315 insertions(+), 338 deletions(-) create mode 100644 lightllm/common/eplb_utils.py create mode 100644 unit_tests/models/deepseek_v4/test_memory_profile.py diff --git a/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py b/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py index 8bcf99b992..13906d0d8a 100644 --- a/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py +++ b/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py @@ -22,7 +22,12 @@ def custom_del(self: torch.Tensor): if hasattr(self, "storage_weak_ptr"): storage_weak_ptr = self.storage_weak_ptr else: - storage_weak_ptr = self.untyped_storage()._weak_ref() + try: + storage_weak_ptr = self.untyped_storage()._weak_ref() + except RuntimeError: + # Some tensor implementations, including UndefinedTensorImpl, + # have no backing storage. Their destructor must stay silent. + return UntypedStorage._free_weak_ref(storage_weak_ptr) if storage_weak_ptr in g_cache_manager.ptr_to_bufnode: g_cache_manager.changed_ptr.add(storage_weak_ptr) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py index a7504e975e..1ea8eb1643 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -33,7 +33,7 @@ def build_logical_to_physical_map( ]: # logical_to_physical [num_logical_experts, num_ranks], replica_counts [num_logical_experts] """构建单层逻辑 expert 到物理副本的映射""" - logical_to_physical, replica_counts = _build_layer_maps( + logical_to_physical, replica_counts = build_logical_to_physical_maps_for_layers( redundant_expert_ids.unsqueeze(0), num_logical_experts, source_rank=source_rank, @@ -51,18 +51,34 @@ def build_logical_to_physical_maps_for_layers( torch.Tensor, # logical_to_physical, shape [num_layers, num_logical_experts, num_ranks] torch.Tensor, # replica_counts, shape [num_layers, num_logical_experts] ]: - """为调用方传入的多个指定层构建逻辑 expert 到物理副本的 CPU int32 映射。 + """Build stable CPU int32 maps for the supplied layers without modifying the input.""" + if redundant_expert_ids_by_layer.ndim != 3: + raise ValueError("redundant_expert_ids_by_layer must be [layers, ranks, num_redundant_experts_per_rank]") + num_ranks, num_redundant_experts_per_rank = redundant_expert_ids_by_layer.shape[1:] + assert num_logical_experts % num_ranks == 0 + layout = _get_physical_expert_layout(num_logical_experts, num_ranks, num_redundant_experts_per_rank) + logical_to_physical, replica_counts = _build_global_replica_maps_for_layers(redundant_expert_ids_by_layer, layout) + if source_rank is None: + return logical_to_physical, replica_counts - 第一维是层,不是请求或 token 的 batch,也不会自动处理模型中的其他层。 - 全局构建阶段第 0 列固定为主副本,其余列按 rank-major、slot-major 的稳定顺序 - 写入冗余副本。指定 source_rank 后,先筛选源节点内副本,再按 source_rank - 轮转候选前缀;最终返回映射的第 0 列只是第一个候选,不保证仍是主副本。 - """ - return _build_layer_maps( - redundant_expert_ids_by_layer, - num_logical_experts, + assert node_world_size is not None + replica_positions = torch.arange(num_ranks, dtype=torch.int64) + compact_maps_by_layer, selected_counts_by_layer = _select_source_node_replicas( + logical_to_physical, + replica_counts, source_rank=source_rank, node_world_size=node_world_size, + num_physical_experts_per_rank=layout.num_physical_experts_per_rank, + replica_positions=replica_positions, + ) + return ( + _rotate_selected_replicas( + compact_maps_by_layer, + selected_counts_by_layer, + source_rank=source_rank, + replica_positions=replica_positions, + ), + selected_counts_by_layer, ) @@ -81,31 +97,19 @@ def select_improving_placements( assert current_placement.shape == candidate_placement.shape current_rank_load = _estimate_rank_load(expert_load, current_placement, expert_alignment, node_world_size) candidate_rank_load = _estimate_rank_load(expert_load, candidate_placement, expert_alignment, node_world_size) - if expert_load.ndim == 2: - current_critical = current_rank_load.max(dim=1).values - candidate_critical = candidate_rank_load.max(dim=1).values - else: - current_critical = current_rank_load.max(dim=2).values.sum(dim=0) - candidate_critical = candidate_rank_load.max(dim=2).values.sum(dim=0) + current_critical = current_rank_load.max(dim=-1).values.sum(dim=0) + candidate_critical = candidate_rank_load.max(dim=-1).values.sum(dim=0) # Each changed layer must reduce its own critical load. All selected # changes must then collectively meet the configured model-level # critical-load reduction threshold, avoiding low-gain migrations. improved = candidate_critical < current_critical selected = current_placement.clone() selected[improved] = candidate_placement[improved] - if current_rank_load.ndim == 2: - selected_rank_load = torch.where(improved[:, None], candidate_rank_load, current_rank_load) - else: - selected_rank_load = torch.where(improved[None, :, None], candidate_rank_load, current_rank_load) + selected_rank_load = torch.where(improved[:, None], candidate_rank_load, current_rank_load) model_current_critical = current_critical.sum() - if expert_load.ndim == 2: - model_current_mean = current_rank_load.mean(dim=1).sum() - model_selected_critical = selected_rank_load.max(dim=1).values.sum() - model_selected_mean = selected_rank_load.mean(dim=1).sum() - else: - model_current_mean = current_rank_load.mean(dim=2).sum() - model_selected_critical = selected_rank_load.max(dim=2).values.sum() - model_selected_mean = selected_rank_load.mean(dim=2).sum() + model_current_mean = current_rank_load.mean(dim=-1).sum() + model_selected_critical = selected_rank_load.max(dim=-1).values.sum() + model_selected_mean = selected_rank_load.mean(dim=-1).sum() model_ratio = model_current_critical / model_current_mean.clamp_min(1.0) candidate_model_ratio = model_selected_critical / model_selected_mean.clamp_min(1.0) candidate_rebalance_gain = (model_current_critical - model_selected_critical) / model_current_critical.clamp_min( @@ -137,31 +141,28 @@ def plan_redundant_experts( current_placement: torch.Tensor | None = None, stickiness: float = 0.0, ) -> torch.Tensor: - """Plan replicas using source-node-local copies, with global fallback. + """Plan replicas from [samples, layers, source_nodes, experts] loads. - With ``current_placement`` and a positive ``stickiness``, a candidate that + With ``current_placement`` and positive ``stickiness``, a candidate that keeps an expert on its current rank receives a bonus of ``stickiness * mean per-layer expert load``. This preserves rank membership, not a particular redundant physical slot; target slots are canonicalized against the current live rows before transfer and metadata publication. A rank membership only changes when the move improves the critical-load objective by more than that margin. - Without them the planning is bit-identical to the legacy behavior. + With zero stickiness, placement is determined solely by the load objective. """ - assert expert_load.ndim in (2, 3, 4) if expert_alignment is not None: assert expert_alignment > 0 - use_legacy_topology_preference = expert_load.ndim < 4 - legacy_node_world_size = node_world_size if use_legacy_topology_preference else None - source_load, _squeeze_sample, node_world_size = _as_source_node_load(expert_load, num_ranks, node_world_size) - num_samples, num_layers, num_nodes, num_logical_experts = source_load.shape + node_world_size = _resolve_node_world_size(expert_load, num_ranks, node_world_size) + num_samples, num_layers, num_nodes, num_logical_experts = expert_load.shape assert num_logical_experts % num_ranks == 0 assert num_redundant_experts_per_rank > 0 num_experts_per_rank = num_logical_experts // num_ranks num_redundant = num_ranks * num_redundant_experts_per_rank assert num_redundant <= num_logical_experts * (num_ranks - 1) - load = source_load.to(dtype=torch.float64, device="cpu") + load = expert_load.to(dtype=torch.float64, device="cpu") placement = torch.full((num_layers, num_ranks, num_redundant_experts_per_rank), -1, dtype=torch.int64) owner_rank = torch.arange(num_logical_experts, dtype=torch.int64) // num_experts_per_rank if current_placement is not None: @@ -182,12 +183,6 @@ def plan_redundant_experts( remaining_slots = torch.full((num_layers, num_ranks), num_redundant_experts_per_rank, dtype=torch.int64) layer_indices = torch.arange(num_layers, dtype=torch.int64) expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) - rank_nodes = ( - torch.arange(num_ranks, dtype=torch.int64) // legacy_node_world_size - if legacy_node_world_size is not None and legacy_node_world_size < num_ranks - else None - ) - # Every iteration fills one slot per layer. Candidate expert evaluation # is vectorized across all layers and logical experts, which keeps large # GLM/Qwen planning comfortably on the CPU fast path. @@ -200,15 +195,6 @@ def plan_redundant_experts( if remaining_slots[layer, target_rank] == 0: continue candidate_legal = (owner_rank != target_rank) & ~locations[layer, :, target_rank] - # Legacy 2D/3D callers have no source-node axis. Retain the - # previous topology preference for that compatibility path; - # node-aware [S,L,N,E] planning uses only the exact load - # objective below. - if rank_nodes is not None: - existing_on_target_node = locations[layer, :, rank_nodes == rank_nodes[target_rank]].any(dim=1) - new_node_legal = candidate_legal & ~existing_on_target_node - if torch.any(new_node_legal): - candidate_legal = new_node_legal if torch.any(candidate_legal): target_ranks[layer] = target_rank legal[layer] = candidate_legal @@ -288,23 +274,7 @@ def _build_global_replica_maps_for_layers( redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] layout: _PhysicalExpertLayout, ) -> Tuple[torch.Tensor, torch.Tensor]: - """为显式传入的多层冗余布局构建全局逻辑 expert id 到物理位置序列[rank, local_slot]的映射。 - - 输出: - logical_to_physical:CPU ``int32`` Tensor,形状为 - ``[num_layers, num_logical_experts, num_ranks]``。第 0 列固定为主 - 副本,后续列依次存放冗余副本,未使用位置为 ``-1``。 - replica_counts:CPU ``int32`` Tensor,形状为 - ``[num_layers, num_logical_experts]``。每个值包含主副本,并表示 - ``logical_to_physical`` 对应行中有效副本连续前缀的长度。 - - “全局映射”记录每层每个逻辑 expert 在所有 rank 上的主副本和冗余副本所对应的 - 物理 expert ID。第 0 列固定为主副本;冗余副本按所在 rank 从小到大、同一 rank - 内按槽位从小到大的顺序写入后续列。输入参数 - ``redundant_expert_ids_by_layer`` 已经指定每个 rank 的每个冗余槽位存放哪个逻辑 - expert,本函数只将该输入转换为顺序确定的映射,满足主副本优先、冗余槽位对应正确且有效副本连续排列的确定性结果。 - - """ + """Build stable global maps with primary copies first and unused slots set to ``-1``.""" num_layers = redundant_expert_ids_by_layer.shape[0] num_logical_experts = layout.num_logical_experts @@ -358,25 +328,7 @@ def _select_source_node_replicas( num_physical_experts_per_rank: int, replica_positions: torch.Tensor, ) -> Tuple[torch.Tensor, torch.Tensor]: - """按来源节点筛选全局候选副本,并将结果压缩为连续前缀。 - - ``logical_to_physical[layer, logical_expert]`` 的前 ``replica_counts`` 个位置 - 是有效副本,尾部 ``-1`` 表示没有副本;正常输入和输出都不会在有效前缀中 - 出现 ``-1``。连续的 ``node_world_size`` 个 rank 构成一个节点,来源节点为 - ``source_rank // node_world_size``。对每个逻辑 expert,若来源节点有副本, - 则只保留该节点的全部副本,远程副本(包括远程主副本)全部排除;否则保留 - 所有全局有效副本作为回退,避免候选为空。筛选可能选中原有效前缀中不连续的 - 位置,因此按原相对顺序将选中副本复制到新映射的连续前缀,其余位置填 ``-1``, - 并返回新的有效数量;不会原地修改输入。此函数只筛选和压缩候选集合,不做 - 负载优化、最终副本选择或按 source rank 轮转,轮转由后续函数完成。 - - 输出: - - ``compact_maps_by_layer``:与 ``logical_to_physical`` 同 shape - ``[num_layers, num_logical_experts, num_ranks]``,dtype/device 相同;每行 - 为筛选后副本的连续有效前缀,尾部为 ``-1``。 - - ``selected_counts_by_layer``:输入同 device 的 ``int32`` Tensor,shape 为 - ``[num_layers, num_logical_experts]``;每个值是对应输出行的有效前缀长度。 - """ + """Keep source-node replicas when available, otherwise fall back to all stable candidates.""" num_layers, num_logical_experts, _max_replicas = logical_to_physical.shape source_node = source_rank // node_world_size output_positions = replica_positions.view(1, 1, -1) @@ -412,105 +364,51 @@ def _rotate_selected_replicas( """根据 source_rank 循环调整每个逻辑 expert 的候选副本顺序,让不同源 rank 优先使用不同副本,同时保持候选副本集合和副本数量不变""" output_positions = replica_positions.view(1, 1, -1) selected_count64_by_layer = selected_count_by_layer.to(torch.int64).unsqueeze(-1) - rotation_by_layer = source_rank % selected_count64_by_layer - source_positions_by_layer = (output_positions + rotation_by_layer) % selected_count64_by_layer + source_positions_by_layer = (output_positions + source_rank) % selected_count64_by_layer maps_by_layer = compact_maps_by_layer.gather(2, source_positions_by_layer) maps_by_layer.masked_fill_(output_positions >= selected_count64_by_layer, -1) return maps_by_layer -def _build_layer_maps( - redundant_expert_ids_by_layer: torch.Tensor, - num_logical_experts: int, # 逻辑 expert 的总数。 - source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 - node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 -) -> Tuple[torch.Tensor, torch.Tensor]: - """构建调用方指定层的映射;全局构建、节点筛选和轮转均在此完成。""" - if redundant_expert_ids_by_layer.ndim != 3: - raise ValueError("redundant_expert_ids_by_layer must be [layers, ranks, num_redundant_experts_per_rank]") - num_ranks, num_redundant_experts_per_rank = redundant_expert_ids_by_layer.shape[1:] - assert num_logical_experts % num_ranks == 0 - layout = _get_physical_expert_layout(num_logical_experts, num_ranks, num_redundant_experts_per_rank) - logical_to_physical, replica_counts = _build_global_replica_maps_for_layers(redundant_expert_ids_by_layer, layout) - if source_rank is None: - return logical_to_physical, replica_counts - - assert node_world_size is not None - replica_positions = torch.arange(num_ranks, dtype=torch.int64) - compact_maps_by_layer, selected_counts_by_layer = _select_source_node_replicas( - logical_to_physical, - replica_counts, - source_rank=source_rank, - node_world_size=node_world_size, - num_physical_experts_per_rank=layout.num_physical_experts_per_rank, - replica_positions=replica_positions, - ) - return ( - _rotate_selected_replicas( - compact_maps_by_layer, - selected_counts_by_layer, - source_rank=source_rank, - replica_positions=replica_positions, - ), - selected_counts_by_layer, - ) - - def _estimate_rank_load( expert_load: torch.Tensor, redundant_expert_ids: torch.Tensor, expert_alignment: int | None = None, node_world_size: int | None = None, ) -> torch.Tensor: - """Estimate runtime source-node-local routing load per physical expert. + """Estimate [samples, layers, ranks] load from source-node-local routing. - ``expert_load`` accepts the historic ``[layers, experts]`` and - ``[samples, layers, experts]`` forms, which are both one source node, and - the distributed ``[samples, layers, source_nodes, experts]`` form. Source - loads are kept separate until they are assigned to physical replicas, then - combined before applying the per-expert alignment used by DeepEP. + Source loads remain separate until assigned to physical replicas, then + combine before the per-expert alignment used by DeepEP. """ - source_load, squeeze_sample, node_world_size = _as_source_node_load( - expert_load, redundant_expert_ids.shape[1], node_world_size - ) - num_samples, num_layers, num_nodes, num_logical_experts = source_load.shape + node_world_size = _resolve_node_world_size(expert_load, redundant_expert_ids.shape[1], node_world_size) + num_samples, num_layers, num_nodes, num_logical_experts = expert_load.shape assert redundant_expert_ids.ndim == 3 and redundant_expert_ids.shape[0] == num_layers num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape[1:] assert num_logical_experts % num_ranks == 0 if expert_alignment is not None: assert expert_alignment > 0 - locations = _expert_locations(redundant_expert_ids, num_logical_experts) - route = _source_route(locations, num_nodes, node_world_size) - physical_load = torch.einsum("slne,lner->sler", source_load.to(torch.float64), route) - if expert_alignment is not None: - physical_load = torch.ceil(physical_load / expert_alignment) * expert_alignment - rank_load = physical_load.sum(dim=2) - return rank_load.squeeze(0) if squeeze_sample else rank_load - - -def _as_source_node_load( - expert_load: torch.Tensor, num_ranks: int, node_world_size: int | None -) -> Tuple[torch.Tensor, bool, int]: - """Normalize load to ``[samples, layers, source_nodes, experts]``.""" - assert expert_load.ndim in (2, 3, 4) - squeeze_sample = expert_load.ndim == 2 - if expert_load.ndim == 2: - source_load = expert_load.unsqueeze(0).unsqueeze(2) - elif expert_load.ndim == 3: - source_load = expert_load.unsqueeze(2) - else: - source_load = expert_load - num_nodes = source_load.shape[2] - # Historic 2D/3D loads represent one source node containing every rank. - if expert_load.ndim < 4: - return source_load, squeeze_sample, num_ranks + rank_load = _expert_rank_load_all( + expert_load, + _expert_locations(redundant_expert_ids, num_logical_experts), + num_nodes, + node_world_size, + expert_alignment, + ).sum(dim=2) + return rank_load + + +def _resolve_node_world_size(expert_load: torch.Tensor, num_ranks: int, node_world_size: int | None) -> int: + """Validate production [samples, layers, source_nodes, experts] planner loads.""" + assert expert_load.ndim == 4 + num_nodes = expert_load.shape[2] if node_world_size is None: assert num_ranks % num_nodes == 0 node_world_size = num_ranks // num_nodes assert 0 < node_world_size <= num_ranks and num_ranks % node_world_size == 0 assert num_nodes == num_ranks // node_world_size - return source_load, squeeze_sample, node_world_size + return node_world_size def _expert_locations(redundant_expert_ids: torch.Tensor, num_logical_experts: int) -> torch.Tensor: diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index c00a35f600..acbc31d261 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -14,24 +14,17 @@ def create_fuse_moe_impl( expert_parallel_state: ExpertParallelState | None = None, ): if expert_parallel_state is not None: - return FuseMoeDeepGEMM( - n_routed_experts=n_routed_experts, - num_fused_shared_experts=num_fused_shared_experts, - routed_scaling_factor=routed_scaling_factor, - quant_method=quant_method, - expert_parallel_state=expert_parallel_state, - ) - - if quant_method.method_name == "awq_marlin": - return FuseMoeMarlin( - n_routed_experts=n_routed_experts, - num_fused_shared_experts=num_fused_shared_experts, - routed_scaling_factor=routed_scaling_factor, - quant_method=quant_method, - ) - return FuseMoeTriton( + impl_cls = FuseMoeDeepGEMM + elif quant_method.method_name == "awq_marlin": + impl_cls = FuseMoeMarlin + else: + impl_cls = FuseMoeTriton + kwargs = dict( n_routed_experts=n_routed_experts, num_fused_shared_experts=num_fused_shared_experts, routed_scaling_factor=routed_scaling_factor, quant_method=quant_method, ) + if expert_parallel_state is not None: + kwargs["expert_parallel_state"] = expert_parallel_state + return impl_cls(**kwargs) 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 f234f09896..7ddd93fbb4 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 @@ -46,13 +46,11 @@ def _select_experts( """选择 expert;EPLB prefill 统一由融合路径返回 physical ID。""" assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" eplb = self.eplb - eplb_active = eplb is not None - if is_prefill is True and eplb_active: + if is_prefill is True and eplb is not None: from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb group_score_topk_num = 2 if topk_group == 4 and num_expert_group == 8 and top_k == 8 else 1 topk_weights, topk_ids, logical_topk_ids = triton_grouped_topk_eplb( - hidden_states=input_tensor, gating_output=router_logits, correction_bias=correction_bias, topk=top_k, @@ -84,8 +82,6 @@ def _select_experts( num_expert_group=num_expert_group, scoring_func=scoring_func, ) - if per_expert_scale is not None: - topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype) origin_topk_ids = topk_ids if self.routed_scaling_factor != 1.0: topk_weights.mul_(self.routed_scaling_factor) @@ -348,9 +344,7 @@ def _primary_weight_pack(self, weight_pack: WeightPack) -> WeightPack: """返回所有 decode 路径使用的缓存本地主副本视图。""" if self.eplb is None: return weight_pack - cache = getattr(self, "_primary_weight_pack_cache", None) - if cache is None: - cache = self._primary_weight_pack_cache = {} + cache = self._primary_weight_pack_cache cache_key = id(weight_pack) primary = cache.get(cache_key) if primary is None: @@ -362,11 +356,6 @@ def _primary_weight_pack(self, weight_pack: WeightPack) -> WeightPack: if weight_pack.weight_scale is not None else None ), - weight_zero_point=( - getattr(weight_pack, "weight_zero_point", None)[:num_primary_experts_per_rank] - if getattr(weight_pack, "weight_zero_point", None) is not None - else None - ), ) cache[cache_key] = primary return primary diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py index 8544b22e32..7651282b2f 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py @@ -440,7 +440,6 @@ def triton_grouped_topk( def triton_grouped_topk_eplb( - hidden_states: torch.Tensor, gating_output: torch.Tensor, correction_bias: torch.Tensor, topk: int, diff --git a/lightllm/common/eplb_utils.py b/lightllm/common/eplb_utils.py new file mode 100644 index 0000000000..63e3dea069 --- /dev/null +++ b/lightllm/common/eplb_utils.py @@ -0,0 +1,16 @@ +"""Small, dependency-free EPLB helpers shared by transfer and model profiling.""" + + +EPLB_MAX_STAGING_DEPTH = 8 + + +def extract_eplb_expert_tensors(weight): + result = [] + for pack_name in ("w13", "w2"): + pack = getattr(weight, pack_name) + for value_name in ("weight", "weight_scale"): + tensor = getattr(pack, value_name, None) + if tensor is not None: + assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" + result.append((f"{pack_name}.{value_name}", tensor)) + return result diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index 9171566599..bc667b7d3a 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -31,13 +31,30 @@ class MemoryManager: operator_class = NormalMemOperator - def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False, mem_fraction=0.9): + def __init__( + self, + size, + dtype, + head_num, + head_dim, + layer_num, + always_copy=False, + mem_fraction=0.9, + memory_reservations=None, + ): self.size = size self.head_num = head_num self.head_dim = head_dim self.layer_num = layer_num self.always_copy = always_copy self.dtype = dtype + # Named reservations are allocations made after KV profiling. They are + # deliberately outside get_fixed_memory_size(): model-specific exact KV + # geometry owns fixed bytes, while these values are deducted once from + # the profile budget only. + self.memory_reservations = dict(memory_reservations or {}) + if any(value < 0 for value in self.memory_reservations.values()): + raise ValueError(f"memory reservations must be non-negative: {self.memory_reservations}") # profile the max total token num if the size is None self.profile_size(mem_fraction) @@ -83,6 +100,13 @@ def get_pd_kv_move_buffer_size(self): ) return math.prod(shape) * torch._utils._element_size(self.dtype) + def get_fixed_memory_size(self): + return self.get_pd_kv_move_buffer_size() + + def get_profiled_size(self, available_memory_bytes): + """Select token capacity after fixed and post-profile reservations.""" + return int(available_memory_bytes / self.get_cell_size()) + def profile_size(self, mem_fraction): if self.size is not None: return @@ -91,16 +115,25 @@ def profile_size(self, mem_fraction): world_size = dist.get_world_size() available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) cell_size = self.get_cell_size() - pd_kv_move_buffer_size = self.get_pd_kv_move_buffer_size() - available_memory_bytes = available_memory * 1024 ** 3 - pd_kv_move_buffer_size - self.size = int(available_memory_bytes / cell_size) + fixed_memory_size = self.get_fixed_memory_size() + reservations = getattr(self, "memory_reservations", {}) + reserved_memory_size = sum(reservations.values()) + available_memory_bytes = available_memory * 1024 ** 3 - fixed_memory_size - reserved_memory_size + if available_memory_bytes <= 0: + raise RuntimeError( + f"{type(self).__name__} fixed buffers require {fixed_memory_size / 1024**3:.2f} GB, " + f"plus {reserved_memory_size / 1024**3:.2f} GB reservations, " + f"but only {available_memory:.2f} GB is available" + ) + self.size = self.get_profiled_size(available_memory_bytes) if world_size > 1: tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}") dist.all_reduce(tensor, op=dist.ReduceOp.MIN) self.size = tensor.item() logger.info( f"{str(available_memory)} GB space is available after load the model weight\n" - f"{str(pd_kv_move_buffer_size / 1024 ** 2)} MB is reserved for PD KV transfer buffer\n" + f"{str(fixed_memory_size / 1024 ** 2)} MB is reserved for fixed KV cache buffers\n" + f"{reservations} bytes are reserved for post-profile model buffers\n" f"{str(cell_size / 1024 ** 2)} MB is the size of one token kv cache\n" f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" ) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 06b02b35ef..582a4470d6 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -265,16 +265,13 @@ def init_model(self, kvargs): self.model.eplb_manager = EPLBManager(self.model) dist.barrier() - self.start_infer_loops() - return - - def start_infer_loops(self): # 启动infer_loop_thread, 启动两个线程进行推理,对于具备双batch推理折叠得场景 # 可以降低 cpu overhead,大幅提升gpu得使用率。 self.infer_loop_thread = threading.Thread(target=self.infer_loop, daemon=True) self.infer_loop_thread.start() self.infer_loop_thread1 = threading.Thread(target=self.infer_loop, daemon=True) self.infer_loop_thread1.start() + return def init_custom(self): pass diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index fbff098acf..456565e81d 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -70,7 +70,6 @@ def __init__(self, model: TpPartBaseModel): # base window before the next fixed sampling boundary. self._continuous_collection_start_step: Optional[int] = None self._continuous_collection_end_step: Optional[int] = self.step_interval - self._sampling_pending = False self._steady_collection_end_step: Optional[int] = None self._reset_recorded_samples() self._set_recording(True) @@ -108,7 +107,6 @@ def step(self): continuous_start = self._continuous_collection_start_step continuous_end = self._continuous_collection_end_step if continuous_end is not None: - # 启动或稀疏样本不足时连续采样,保证负载统计可靠。 if continuous_start is not None and self.prefill_steps == continuous_start: self._set_recording(True) if self.prefill_steps >= continuous_end: @@ -116,19 +114,17 @@ def step(self): return sampling_interval = self.sampling_interval phase = self.prefill_steps % sampling_interval - # _sampling_pending=True 表示稳态采样窗口已启动,防止重复启动;到期后清除标记并评估。 - if self._sampling_pending: - steady_collection_end_step = self._steady_collection_end_step - if steady_collection_end_step is None or self.prefill_steps >= steady_collection_end_step: - self._clear_steady_collection() + steady_collection_end_step = self._steady_collection_end_step + if steady_collection_end_step is not None: + if self.prefill_steps >= steady_collection_end_step: + self._steady_collection_end_step = None self._start_evaluation() return if sampling_interval == 1: self._start_evaluation() return if phase == sampling_interval - self._steady_sample_window_steps(): - # 稳态仅在周期末采样少量 step,降低路由计数和评估开销。 - self._start_steady_sampling_window(self.prefill_steps + self._steady_sample_window_steps()) + self._arm_steady_collection(self.prefill_steps + self._steady_sample_window_steps()) def _set_recording(self, enabled: bool): for state in self._eplb_states: @@ -149,17 +145,12 @@ def _clear_continuous_collection(self): self._continuous_collection_start_step = None self._continuous_collection_end_step = None - def _clear_steady_collection(self): - self._sampling_pending = False - self._steady_collection_end_step = None - def _steady_sample_window_steps(self) -> int: return min(EPLB_STEADY_SAMPLE_STEPS, self.sampling_interval) - def _start_steady_sampling_window(self, collection_end_step: int): + def _arm_steady_collection(self, collection_end_step: int): """Start the fixed sparse window without moving its evaluation boundary.""" self._reset_recorded_samples() - self._sampling_pending = True self._steady_collection_end_step = collection_end_step self._set_recording(True) @@ -167,7 +158,7 @@ def _begin_continuous_collection(self): minimum_end = self.prefill_steps + self.step_interval collection_end = -(-minimum_end // self.sampling_interval) * self.sampling_interval self._reset_recorded_samples() - self._clear_steady_collection() + self._steady_collection_end_step = None self._continuous_collection_start_step = collection_end - self.step_interval self._continuous_collection_end_step = collection_end self._set_recording(self._continuous_collection_start_step == self.prefill_steps) @@ -175,7 +166,7 @@ def _begin_continuous_collection(self): def _prepare_next_sampling_window(self): """Clear the current window and arm the next sparse sampling window.""" self._clear_continuous_collection() - self._clear_steady_collection() + self._steady_collection_end_step = None if self.sampling_interval == 1: self._reset_recorded_samples() self._set_recording(True) @@ -183,7 +174,7 @@ def _prepare_next_sampling_window(self): # There is no later pre-boundary manager step at which to arm a # full clamped window, so arm immediately but keep the same next # fixed boundary. - self._start_steady_sampling_window(self.prefill_steps + self.sampling_interval) + self._arm_steady_collection(self.prefill_steps + self.sampling_interval) else: self._reset_recorded_samples() self._set_recording(False) @@ -533,14 +524,10 @@ def _start_rebalance(self, result): def _imbalance_summary(rank_load: torch.Tensor) -> Dict[str, float]: - if rank_load.ndim == 2: - critical = rank_load.max(dim=1).values - mean = rank_load.mean(dim=1) - elif rank_load.ndim == 3: - critical = rank_load.max(dim=2).values.sum(dim=0) - mean = rank_load.mean(dim=2).sum(dim=0) - else: - raise ValueError("rank_load must be [layers, ranks] or [samples, layers, ranks]") + if rank_load.ndim != 3: + raise ValueError("rank_load must be [samples, layers, ranks]") + critical = rank_load.max(dim=2).values.sum(dim=0) + mean = rank_load.mean(dim=2).sum(dim=0) layer_imbalance = critical / mean.clamp_min(1.0) sorted_imbalance = torch.sort(layer_imbalance).values p95_index = max(0, (95 * layer_imbalance.numel() + 99) // 100 - 1) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index 7e5a553571..e5737fcb7a 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -11,6 +11,8 @@ import torch import torch.distributed as dist +from lightllm.common.eplb_utils import EPLB_MAX_STAGING_DEPTH, extract_eplb_expert_tensors + @dataclass(frozen=True) class TransferStep: @@ -97,8 +99,6 @@ def build_transfer_plan( class _EPLBTransferBase: """Shared live/staging buffers and publish/commit lifecycle.""" - staging_depth = 1 - def __init__(self, weights, transfer_group, global_rank, world_size): self._eplb_states = [weight.expert_parallel_state.eplb for weight in weights] self.transfer_group = transfer_group @@ -106,8 +106,8 @@ def __init__(self, weights, transfer_group, global_rank, world_size): self.world_size = world_size self.num_experts_per_rank = weights[0].expert_parallel_state.num_primary_experts_per_rank self.device = weights[0].w13.weight.device - self.live = [_extract_expert_tensors(weight) for weight in weights] - self._validate_live_layout(weights) + self.live = [extract_eplb_expert_tensors(weight) for weight in weights] + self._validate_live_layout() num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank self.staging = [ [ @@ -135,7 +135,7 @@ def __init__(self, weights, transfer_group, global_rank, world_size): self._thread = None self._needs_staging_reuse_barrier = False - def _validate_live_layout(self, weights) -> None: + def _validate_live_layout(self) -> None: reference = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in self.live[0]] num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank for layer_index, (state, tensors) in enumerate(zip(self._eplb_states, self.live)): @@ -145,9 +145,6 @@ def _validate_live_layout(self, weights) -> None: state.num_redundant_experts_per_rank == num_redundant_slots_per_rank ), "EPLB redundant slot count must match" - def _copy_batch(self, batch, prepared_batch) -> None: - raise NotImplementedError - def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): return [ [ @@ -159,20 +156,9 @@ def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]] for batch_start in range(0, len(layer_plans), self.staging_depth) ] - def prepare_transfer(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): - return [(batch, None) for batch in self._make_batches(layer_plans)] - - def _start_transfer_generation(self) -> None: - """Prepare backend state after the in-flight worker check succeeds.""" - - def _finish_transfer_generation(self) -> None: - """Release backend state only after the migration worker has joined.""" - - def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepared_batches=None) -> None: + def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepared_batches) -> None: if self._thread is not None and self._thread.is_alive(): raise RuntimeError("EPLB transfer is already in flight") - if prepared_batches is None: - prepared_batches = self.prepare_transfer(layer_plans) expected_batch_count = (len(layer_plans) + self.staging_depth - 1) // self.staging_depth if len(prepared_batches) != expected_batch_count: raise ValueError("EPLB prepared batch count does not match layer-plan batches") @@ -263,7 +249,7 @@ class _PreparedBatch: def __init__(self, weights, transfer_group, global_rank, world_size): # Reuse at most eight layer buffers to bound EPLB staging memory. - self.staging_depth = min(8, len(weights)) + self.staging_depth = min(EPLB_MAX_STAGING_DEPTH, len(weights)) super().__init__(weights, transfer_group, global_rank, world_size) self._nixl_agent = None self._registered_descs = None @@ -619,18 +605,6 @@ def __del__(self): pass -def _extract_expert_tensors(weight) -> List[Tuple[str, torch.Tensor]]: - result = [] - for pack_name in ("w13", "w2"): - pack = getattr(weight, pack_name) - for value_name in ("weight", "weight_scale", "weight_zero_point"): - tensor = getattr(pack, value_name, None) - if tensor is not None: - assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" - result.append((f"{pack_name}.{value_name}", tensor)) - return result - - def _commit_staging_rows( live: torch.Tensor, staging: torch.Tensor, diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 505588f4ec..c301f01949 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -127,6 +127,7 @@ def get_eplb_placement_stickiness() -> float: return value +@lru_cache(maxsize=None) def get_triton_autotune_level(): return int(os.getenv("LIGHTLLM_TRITON_AUTOTUNE_LEVEL", 0)) diff --git a/lightllm/utils/profile_max_tokens.py b/lightllm/utils/profile_max_tokens.py index 5ec2a9145e..73f172d908 100644 --- a/lightllm/utils/profile_max_tokens.py +++ b/lightllm/utils/profile_max_tokens.py @@ -74,6 +74,14 @@ def profile_mtp_weight_memory(model): weight_memory_before = torch.cuda.memory_allocated() yield target_weight_bytes = torch.cuda.memory_allocated() - weight_memory_before + # Draft construction can opt out of portions of the target's weight layout + # which it never instantiates (for example EPLB redundant rows). + excluded_weight_bytes = int(getattr(model, "get_mtp_profile_weight_exclusion", lambda: 0)()) + if not 0 <= excluded_weight_bytes <= target_weight_bytes: + raise ValueError( + f"invalid MTP profile exclusion {excluded_weight_bytes}; measured target weights={target_weight_bytes}" + ) + target_weight_bytes -= excluded_weight_bytes model.mem_fraction = get_mtp_adjusted_mem_fraction( mem_fraction=model.mem_fraction, target_weight_bytes=target_weight_bytes, diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index e34995a53f..fd656e07c3 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -48,11 +48,11 @@ disable_eplb_model_init, is_eplb_model_init_disabled, ) +from lightllm.common.eplb_utils import extract_eplb_expert_tensors from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( TransferStep, _CudaBatchMemcpy, _commit_staging_rows, - _extract_expert_tensors, align_target_placement, build_transfer_plan, ) @@ -124,6 +124,7 @@ def _validated_expert_parallel_state( def _set_expert_parallel_state(impl, state): impl.expert_parallel_state = state impl.eplb = state.eplb + impl._primary_weight_pack_cache = {} def _manual_runtime_rank_load(source_load, placement, node_world_size, alignment): @@ -211,9 +212,10 @@ def test_parallel_state_derives_expert_layout(): assert state.num_total_physical_experts == 6 -def test_factory_selects_all_paths_and_requires_ep_state(): +def test_factory_selects_all_paths_and_requires_ep_state(monkeypatch): plain_quant = SimpleNamespace(method_name="none") marlin_quant = SimpleNamespace(method_name="awq_marlin") + monkeypatch.setattr(FuseMoeMarlin, "create_workspace", lambda self: None) state = _validated_expert_parallel_state(eplb=False) ep_impl = create_fuse_moe_impl( n_routed_experts=4, @@ -301,11 +303,15 @@ def test_build_initial_redundant_expert_ids( def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): - expert_load = torch.tensor( - [ - [100, 90, 80, 70, 60, 50, 40, 30], - [30, 40, 50, 60, 70, 80, 90, 100], - ] + expert_load = ( + torch.tensor( + [ + [100, 90, 80, 70, 60, 50, 40, 30], + [30, 40, 50, 60, 70, 80, 90, 100], + ] + ) + .unsqueeze(0) + .unsqueeze(2) ) placement = plan_redundant_experts(expert_load, num_ranks=4, num_redundant_experts_per_rank=2) @@ -316,7 +322,7 @@ def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): def test_plan_redundant_experts_minimizes_samplewise_aligned_critical_load(): - samples = torch.tensor([[[300, 20, 20, 200]], [[100, 300, 40, 160]]]) + samples = torch.tensor([[[300, 20, 20, 200]], [[100, 300, 40, 160]]]).unsqueeze(2) placement = plan_redundant_experts(samples, num_ranks=2, num_redundant_experts_per_rank=1, expert_alignment=128) candidates = [torch.tensor([[[left], [right]]]) for left in (2, 3) for right in (0, 1)] @@ -328,7 +334,7 @@ def critical(candidate): def test_select_improving_placements_rejects_regressing_layer(): - expert_load = torch.tensor([[8649, 5740, 5002, 3441]]) + expert_load = torch.tensor([[8649, 5740, 5002, 3441]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) regressing_candidate = torch.tensor([[[1], [0]]]) @@ -348,7 +354,7 @@ def test_select_improving_placements_rejects_regressing_layer(): def test_select_improving_placements_rejects_near_balance_when_gain_is_below_threshold(): - expert_load = torch.tensor([[1, 2, 1, 17]]) + expert_load = torch.tensor([[1, 2, 1, 17]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[3], [0]]]) candidate = torch.tensor([[[3], [1]]]) @@ -367,7 +373,7 @@ def test_select_improving_placements_rejects_near_balance_when_gain_is_below_thr def test_select_improving_placements_accepts_alignment_aware_gain_even_when_current_ranks_are_balanced(): - expert_load = torch.tensor([[100, 129, 100, 129]]) + expert_load = torch.tensor([[100, 129, 100, 129]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[3], [1]]]) @@ -386,7 +392,7 @@ def test_select_improving_placements_accepts_alignment_aware_gain_even_when_curr def test_select_improving_placements_rejects_insufficient_rebalance_gain(): - expert_load = torch.tensor([[1, 1, 6, 7]]) + expert_load = torch.tensor([[1, 1, 6, 7]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[3], [0]]]) @@ -420,7 +426,7 @@ def test_select_improving_placements_rejects_insufficient_rebalance_gain(): def test_select_improving_placements_rejects_invalid_rebalance_gain_threshold( rebalance_gain_threshold, ): - expert_load = torch.tensor([[1, 1, 1, 2]]) + expert_load = torch.tensor([[1, 1, 1, 2]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[3], [0]]]) @@ -434,7 +440,7 @@ def test_select_improving_placements_rejects_invalid_rebalance_gain_threshold( def test_select_improving_placements_accepts_sufficient_rebalance_gain(): - expert_load = torch.tensor([[1, 1, 1, 2]]) + expert_load = torch.tensor([[1, 1, 1, 2]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[3], [0]]]) @@ -455,7 +461,7 @@ def test_select_improving_placements_accepts_sufficient_rebalance_gain(): def test_select_improving_placements_rejects_raw_improvement_that_does_not_improve_aligned_compute(): - expert_load = torch.tensor([[1, 1, 1, 8]]) + expert_load = torch.tensor([[1, 1, 1, 8]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) raw_improving_candidate = torch.tensor([[[3], [0]]]) @@ -476,15 +482,15 @@ def test_select_improving_placements_rejects_raw_improvement_that_does_not_impro def test_estimate_rank_load_aligns_each_sample_before_accumulation(): - samples = torch.tensor([[[20, 0, 0, 0]], [[20, 0, 0, 0]]]) + samples = torch.tensor([[[20, 0, 0, 0]], [[20, 0, 0, 0]]]).unsqueeze(2) placement = torch.tensor([[[2], [0]]]) per_sample = _estimate_rank_load(samples, placement, expert_alignment=128) - accumulated = _estimate_rank_load(samples.sum(dim=0), placement, expert_alignment=128) + accumulated = _estimate_rank_load(samples.sum(dim=0, keepdim=True), placement, expert_alignment=128) assert torch.equal(per_sample[:, 0], torch.tensor([[128.0, 128.0], [128.0, 128.0]])) assert torch.equal(per_sample.sum(dim=0)[0], torch.tensor([256.0, 256.0])) - assert torch.equal(accumulated[0], torch.tensor([128.0, 128.0])) + assert torch.equal(accumulated[0, 0], torch.tensor([128.0, 128.0])) def test_select_improving_placements_rejects_lower_ratio_when_critical_is_unchanged(): @@ -494,7 +500,7 @@ def test_select_improving_placements_rejects_lower_ratio_when_critical_is_unchan [[172, 278, 51, 238]], [[249, 291, 284, 183]], ] - ) + ).unsqueeze(2) current = torch.tensor([[[2], [0]]]) mean_inflating_candidate = torch.tensor([[[2], [1]]]) @@ -523,7 +529,7 @@ def test_select_improving_placements_accepts_five_percent_critical_reduction(): [[287, 175, 236, 179]], [[316, 99, 266, 353]], ] - ) + ).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[2], [1]]]) @@ -538,7 +544,7 @@ def test_select_improving_placements_accepts_five_percent_critical_reduction(): def test_select_improving_placements_rejects_single_layer_gain_below_model_threshold(): # Layer 0 becomes better, but layer 1 dominates model critical load. The # aggregate estimated critical-load reduction gain is below 5%, so neither layer may be changed. - expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 10]]) + expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 10]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]], [[2], [0]]]) candidate = torch.tensor([[[3], [0]], [[2], [0]]]) @@ -553,7 +559,7 @@ def test_select_improving_placements_rejects_single_layer_gain_below_model_thres def test_select_improving_placements_accepts_only_when_model_gain_reaches_threshold(): - expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 5]]) + expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 5]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]], [[2], [0]]]) candidate = torch.tensor([[[3], [0]], [[2], [0]]]) @@ -672,22 +678,23 @@ def test_logical_to_physical_maps_for_layers_match_single_layer_api(source_rank, assert torch.all(maps_by_layer[~valid] == -1) -def test_plan_redundant_experts_prefers_first_replica_on_new_node(): - # One redundant slot per rank leaves legal alternatives on both nodes; - # topology preference therefore puts every first replica away from its - # primary node before considering same-node duplicates. +def test_plan_redundant_experts_prefers_local_node_load_relief(): + source_load = torch.zeros((1, 1, 2, 8), dtype=torch.int64) + source_load[0, 0, 0, 0] = 1024 placement = plan_redundant_experts( - torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]), + source_load, num_ranks=4, num_redundant_experts_per_rank=1, + expert_alignment=128, node_world_size=2, ) - for rank, expert in enumerate(placement[0, :, 0].tolist()): - assert expert // 2 // 2 != rank // 2 + assert placement[0, 1, 0] == 0 + predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) + assert torch.equal(predicted, _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128)) def test_plan_redundant_experts_single_node_matches_default_behavior(): - load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]) + load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]).unsqueeze(0).unsqueeze(2) default = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1) single_node = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1, node_world_size=4) assert torch.equal(single_node, default) @@ -703,7 +710,7 @@ def test_source_node_estimate_matches_local_first_runtime_replica_sharing(): predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) runtime = _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128) - collapsed_global = _estimate_rank_load(source_load.sum(dim=2), placement, expert_alignment=128) + collapsed_global = _estimate_rank_load(source_load.sum(dim=2, keepdim=True), placement, expert_alignment=128) assert torch.equal(predicted, runtime) assert torch.equal(predicted[0, 0], torch.tensor([256.0, 0.0, 128.0, 0.0])) @@ -781,7 +788,7 @@ def _count_moved_slots(current: torch.Tensor, target: torch.Tensor) -> int: def test_sticky_plan_reproduces_current_when_load_unchanged(): generator = torch.Generator().manual_seed(7) - load = torch.randint(1, 1000, (3, 16, 32), generator=generator) + load = torch.randint(1, 1000, (3, 16, 32), generator=generator).unsqueeze(2) placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) replanned = plan_redundant_experts( @@ -799,9 +806,9 @@ def test_sticky_plan_reproduces_current_when_load_unchanged(): def test_sticky_plan_bounded_moves_under_small_perturbation(): generator = torch.Generator().manual_seed(11) - load = torch.randint(100, 1000, (4, 16, 32), generator=generator) + load = torch.randint(100, 1000, (4, 16, 32), generator=generator).unsqueeze(2) placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - noise = torch.rand((4, 16, 32), generator=generator) * 0.1 + 0.95 + noise = torch.rand((4, 16, 32), generator=generator).unsqueeze(2) * 0.1 + 0.95 perturbed = (load.double() * noise).round().to(torch.int64) sticky = plan_redundant_experts(perturbed, 4, 2, current_placement=placement, stickiness=0.1) @@ -826,6 +833,8 @@ def test_sticky_plan_still_churns_under_phase_shift(): for layer in range(layers): before[layer, (4 * layer + offsets) % experts] = 5000 after[layer, (4 * layer + 16 + offsets) % experts] = 5000 + before = before.unsqueeze(0).unsqueeze(2) + after = after.unsqueeze(0).unsqueeze(2) placement = plan_redundant_experts(before, num_ranks=4, num_redundant_experts_per_rank=2) replanned = plan_redundant_experts( @@ -900,13 +909,17 @@ def test_plan_and_broadcast_publishes_canonical_placement(monkeypatch): broadcasts = [] def fixed_selector(*_args, **_kwargs): - rank_load = torch.full((1, 4), 100.0) + rank_load = torch.full((1, 1, 4), 100.0) return candidate.clone(), torch.tensor([True]), {}, rank_load, rank_load def record_broadcast(result_list, **_kwargs): broadcasts.append(result_list[0]) - monkeypatch.setattr(manager_module, "plan_redundant_experts", lambda *_args, **_kwargs: candidate.clone()) + monkeypatch.setattr( + manager_module, + "plan_redundant_experts", + lambda *_args, **_kwargs: candidate.clone(), + ) monkeypatch.setattr(manager_module, "select_improving_placements", fixed_selector) monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) @@ -916,10 +929,10 @@ def record_broadcast(result_list, **_kwargs): assert broadcasts and torch.equal(broadcasts[0]["placement"], manager.current_placement) -def test_stickiness_zero_matches_legacy(): +def test_stickiness_zero_matches_unbiased_plan(): generator = torch.Generator().manual_seed(17) - load = torch.randint(1, 1000, (2, 8, 16), generator=generator) - legacy = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) + load = torch.randint(1, 1000, (2, 8, 16), generator=generator).unsqueeze(2) + unbiased = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) unrelated = build_initial_redundant_expert_ids(16, 4, 2).unsqueeze(0).expand(8, -1, -1).clone() replanned = plan_redundant_experts( @@ -930,10 +943,12 @@ def test_stickiness_zero_matches_legacy(): stickiness=0.0, ) - assert torch.equal(replanned, legacy) + assert torch.equal(replanned, unbiased) -def test_plan_and_broadcast_propagates_rank_zero_error_after_existing_broadcast(monkeypatch): +def test_plan_and_broadcast_propagates_rank_zero_error_after_existing_broadcast( + monkeypatch, +): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.global_rank = 0 manager.world_size = 2 @@ -992,7 +1007,6 @@ def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_eval manager.num_logical_experts = 4 manager.global_rank = 1 manager.evaluation_in_flight = False - manager._sampling_pending = False manager._steady_collection_end_step = None manager._continuous_collection_start_step = None manager._continuous_collection_end_step = None @@ -1011,7 +1025,6 @@ def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_eval assert manager.prefill_steps == 16 assert recordings == [True] assert resets == [True] - assert manager._sampling_pending assert manager._steady_collection_end_step == 20 for _ in range(3): @@ -1022,18 +1035,19 @@ def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_eval manager.step() assert manager.prefill_steps == 20 - assert not manager._sampling_pending + assert manager._steady_collection_end_step is None assert started == [True] -def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary(monkeypatch): +def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary( + monkeypatch, +): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.in_flight = False manager.evaluation_in_flight = False manager.prefill_steps = 0 manager.step_interval = 20 manager.sampling_interval = 3 - manager._sampling_pending = False manager._steady_collection_end_step = None manager._continuous_collection_start_step = None manager._continuous_collection_end_step = None @@ -1046,7 +1060,6 @@ def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary assert resets == [0] assert recordings == [(0, True)] assert manager._steady_collection_end_step == 3 - assert manager._sampling_pending manager.step() manager.step() @@ -1143,7 +1156,7 @@ def join(self): manager.prefill_steps = 1 manager.step_interval = 1 manager.sampling_interval = 1 - manager._sampling_pending = False + manager._steady_collection_end_step = None manager.evaluation_in_flight = True manager._evaluation_lock = threading.Lock() manager._evaluation_error = None @@ -1168,7 +1181,7 @@ def join(self): # The no-improvement backoff changes interval 1 to 4. The clamped # steady window arms immediately but still waits for boundary step 5. assert recordings == [True] - assert manager._sampling_pending + assert manager._steady_collection_end_step is not None assert manager._steady_collection_end_step == 5 assert starts == [] assert manager.prefill_steps == 1 @@ -1997,19 +2010,17 @@ def test_transfer_plan_cross_node_and_stable_source_load_tie_break(): ] -def test_extract_expert_tensors_includes_weight_scale_and_zero_point_in_order(): +def test_extract_expert_tensors_includes_weight_and_scale_in_order(): class Pack: - def __init__(self, offset, scale=True, zero=True): + def __init__(self, offset, scale=True): self.weight = torch.full((3, 2), offset) self.weight_scale = torch.full((3, 1), offset + 1) if scale else None - self.weight_zero_point = torch.full((3, 1), offset + 2) if zero else None - weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False, zero=False)})() - tensors = _extract_expert_tensors(weight) + weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False)})() + tensors = extract_eplb_expert_tensors(weight) assert [name for name, _ in tensors] == [ "w13.weight", "w13.weight_scale", - "w13.weight_zero_point", "w2.weight", ] @@ -2352,6 +2363,8 @@ def synchronize(self): transfer._error = None transfer._thread = None transfer._needs_staging_reuse_barrier = True + transfer._start_transfer_generation = lambda: None + transfer._finish_transfer_generation = lambda: None transfer.transfer_group = "transfer-group" copied = [] @@ -2371,7 +2384,7 @@ def copy_batch(batch, _prepared_batch): ) plans = [(0, []), (1, []), (2, [])] - prepared_batches = transfer.prepare_transfer(plans) + prepared_batches = [(batch, None) for batch in transfer._make_batches(plans)] monkeypatch.setattr(transfer, "_make_batches", lambda _plans: pytest.fail("start must reuse prepared batches")) transfer.start(plans, prepared_batches) deadline = time.monotonic() + 2 @@ -2399,7 +2412,9 @@ def copy_batch(batch, _prepared_batch): transfer.finish() -def test_transfer_finalization_failure_stays_in_worker_and_success_finalizes_once(monkeypatch): +def test_transfer_finalization_failure_stays_in_worker_and_success_finalizes_once( + monkeypatch, +): def make_transfer(finalize): transfer = object.__new__(transfer_module._EPLBTransferBase) transfer.backend = "test" @@ -2418,6 +2433,7 @@ def make_transfer(finalize): transfer._thread = None transfer._needs_staging_reuse_barrier = False transfer._copy_batch = lambda _batch, _prepared_batch: None + transfer._start_transfer_generation = lambda: None transfer._finish_transfer_generation = finalize return transfer @@ -2425,13 +2441,13 @@ def make_transfer(finalize): finalized_before_publish = [] success = make_transfer(lambda: finalized_before_publish.append(len(success._pending))) - success.start([(0, [])]) + success.start([(0, [])], [([(0, [], 0, [])], None)]) success.finish() assert finalized_before_publish == [0] assert success.pending_layers() == [(0, 0)] failed = make_transfer(lambda: (_ for _ in ()).throw(RuntimeError("cache boom"))) - failed.start([(0, [])]) + failed.start([(0, [])], [([(0, [], 0, [])], None)]) failed._thread.join() assert list(failed._pending) == [] with pytest.raises(RuntimeError, match="EPLB migration worker failed") as exc_info: @@ -2448,7 +2464,7 @@ def test_manager_rearms_after_rebalance_for_interval_one(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.step_interval = 1 manager.sampling_interval = 1 - manager._sampling_pending = False + manager._steady_collection_end_step = None manager._continuous_collection_start_step = 0 manager.weights = [] manager._eplb_states = [] @@ -2459,7 +2475,7 @@ def test_manager_rearms_after_rebalance_for_interval_one(): manager._finish_rebalance() assert manager.in_flight is False assert recording_calls == [True] - assert not manager._sampling_pending + assert manager._steady_collection_end_step is None assert manager._continuous_collection_start_step is None @@ -2563,7 +2579,7 @@ def test_begin_continuous_collection_uses_full_window_at_fixed_boundary(monkeypa manager.prefill_steps = 36 manager.step_interval = 20 manager.sampling_interval = 20 - manager._sampling_pending = True + manager._steady_collection_end_step = manager.prefill_steps + 1 recordings, resets, starts = [], [], [] manager._set_recording = lambda enabled: recordings.append(enabled) manager._reset_recorded_samples = lambda: resets.append(True) @@ -2576,7 +2592,7 @@ def test_begin_continuous_collection_uses_full_window_at_fixed_boundary(monkeypa assert manager._continuous_collection_end_step == 60 assert recordings == [False] assert resets == [True] - assert not manager._sampling_pending + assert manager._steady_collection_end_step is None for _ in range(4): manager.step() @@ -2595,7 +2611,7 @@ def test_begin_continuous_collection_preserves_full_window_at_sparse_boundary(): manager.prefill_steps = 80 manager.step_interval = 20 manager.sampling_interval = 80 - manager._sampling_pending = False + manager._steady_collection_end_step = None recordings = [] manager._set_recording = lambda enabled: recordings.append(enabled) manager._reset_recorded_samples = lambda: None @@ -2646,7 +2662,7 @@ def test_continuous_collection_evaluates_only_after_one_full_base_window(monkeyp manager.prefill_steps = 0 manager.step_interval = 20 manager.sampling_interval = 320 - manager._sampling_pending = False + manager._steady_collection_end_step = None manager.evaluation_in_flight = False started = [] monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) @@ -2696,7 +2712,7 @@ def test_sparse_backoff_arms_and_evaluates_only_at_new_interval_boundary(monkeyp manager.prefill_steps = 18 manager.step_interval = 20 manager.sampling_interval = 80 - manager._sampling_pending = False + manager._steady_collection_end_step = None manager._continuous_collection_start_step = None manager._continuous_collection_end_step = None manager.evaluation_in_flight = False @@ -2721,13 +2737,13 @@ def test_sparse_backoff_arms_and_evaluates_only_at_new_interval_boundary(monkeyp assert manager.prefill_steps == 76 assert recordings == [True] assert resets == [True] - assert manager._sampling_pending + assert manager._steady_collection_end_step is not None for _ in range(4): manager.step() assert manager.prefill_steps == 80 assert starts == [True] - assert not manager._sampling_pending + assert manager._steady_collection_end_step is None def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): @@ -2785,7 +2801,7 @@ def test_first_rebalance_completion_switches_to_four_step_sparse_window(monkeypa manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.step_interval = 20 manager.sampling_interval = manager.step_interval - manager._sampling_pending = False + manager._steady_collection_end_step = None manager.weights = [] manager._eplb_states = [] manager.target_placement = torch.zeros(1) @@ -2891,12 +2907,12 @@ def enqueue(self, descriptor, stream): [ ("w13.weight", 1000, 32), ("w13.weight_scale", 2000, 32), - ("w13.weight_zero_point", 3000, 32), + ("w2.weight", 3000, 32), ] ] transfer._push_staging_row_layout = { - 1: [[("w13.weight", 4000, 32), ("w13.weight_scale", 5000, 32), ("w13.weight_zero_point", 6000, 32)]], - 2: [[("w13.weight", 7000, 32), ("w13.weight_scale", 8000, 32), ("w13.weight_zero_point", 9000, 32)]], + 1: [[("w13.weight", 4000, 32), ("w13.weight_scale", 5000, 32), ("w2.weight", 6000, 32)]], + 2: [[("w13.weight", 7000, 32), ("w13.weight_scale", 8000, 32), ("w2.weight", 9000, 32)]], } transfer._get_remote_read = lambda *_args: None transfer._wait_xfers = lambda _xfers: None @@ -3529,7 +3545,6 @@ def test_grouped_topk_eplb_matches_topk_mapping_and_counting(record_load, tokens torch.ones(logical_ids.numel(), dtype=torch.int64, device="cuda"), ) fused_weights, fused_ids, fused_logical_ids = triton_grouped_topk_eplb( - hidden_states, gating_output, correction_bias, topk, @@ -3599,7 +3614,6 @@ def test_global_topk_eplb_supports_logical_ids_and_counting(record_load, tokens) expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True) weights, physical_ids, logical_ids = triton_grouped_topk_eplb( - hidden_states=torch.empty((tokens, 1), dtype=torch.float32, device="cuda"), gating_output=gating_output, correction_bias=torch.randn((experts,), dtype=torch.float32, device="cuda"), topk=topk, @@ -3630,7 +3644,6 @@ def test_triton_grouped_topk_eplb_empty_tokens_skips_kernel(): experts = 64 counter = torch.zeros((1, experts), dtype=torch.int64, device="cuda") weights, physical_ids, logical_ids = triton_grouped_topk_eplb( - hidden_states=torch.empty((0, 1), device="cuda"), gating_output=torch.empty((0, experts), device="cuda"), correction_bias=None, topk=4, diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 46c218ad45..3eff70f940 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -90,7 +90,7 @@ def _run_layers( callback=lambda layer_index: None, before_commit_callback=lambda layer_index: None, ): - transfer.start(layer_plans) + transfer.start(layer_plans, transfer.prepare_transfer(layer_plans)) committed = 0 while committed < len(layer_plans): pending = _wait_for_ready_prefix(transfer, control_group) diff --git a/unit_tests/models/deepseek_v4/test_memory_profile.py b/unit_tests/models/deepseek_v4/test_memory_profile.py new file mode 100644 index 0000000000..4c25586831 --- /dev/null +++ b/unit_tests/models/deepseek_v4/test_memory_profile.py @@ -0,0 +1,64 @@ +from types import SimpleNamespace + +import pytest + +from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager +from lightllm.utils import profile_max_tokens + + +@pytest.mark.parametrize("exclusion,expected", [(0, 1000), (200, 800), (None, 1000)]) +def test_mtp_profile_exclusion_adjustment(monkeypatch, exclusion, expected): + seen = [] + values = iter((100, 1100)) + monkeypatch.setattr(profile_max_tokens.torch.cuda, "memory_allocated", lambda: next(values)) + monkeypatch.setattr(profile_max_tokens, "get_mtp_weight_layer_num", lambda: 1) + monkeypatch.setattr( + profile_max_tokens, "get_mtp_adjusted_mem_fraction", lambda **kw: seen.append(kw["target_weight_bytes"]) or 0.5 + ) + attrs = dict( + max_total_token_num=None, + is_mtp_draft_model=False, + args=SimpleNamespace(mtp_mode="x"), + config={"n_layer": 1}, + mem_fraction=0.8, + ) + if exclusion is not None: + attrs["get_mtp_profile_weight_exclusion"] = lambda: exclusion + model = SimpleNamespace(**attrs) + with profile_max_tokens.profile_mtp_weight_memory(model): + pass + assert seen == [expected] + + +@pytest.mark.parametrize("exclusion", [-1, 1001]) +def test_mtp_profile_exclusion_validation(monkeypatch, exclusion): + values = iter((100, 1100)) + monkeypatch.setattr(profile_max_tokens.torch.cuda, "memory_allocated", lambda: next(values)) + model = SimpleNamespace( + max_total_token_num=None, + is_mtp_draft_model=False, + args=SimpleNamespace(mtp_mode="x"), + config={"n_layer": 1}, + mem_fraction=0.8, + get_mtp_profile_weight_exclusion=lambda: exclusion, + ) + with pytest.raises(ValueError, match="invalid MTP profile exclusion"): + with profile_max_tokens.profile_mtp_weight_memory(model): + pass + + +@pytest.mark.parametrize("reservations,expected", [({}, 252), ({"x": 20}, 247)]) +def test_memory_manager_profile_reservation_once(monkeypatch, reservations, expected): + monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.torch.cuda.empty_cache", lambda: None) + monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.dist.get_world_size", lambda: 1) + monkeypatch.setattr( + "lightllm.common.kv_cache_mem_manager.mem_manager.get_available_gpu_memory", lambda w: 1024 / 1024 ** 3 + ) + monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.get_total_gpu_memory", lambda: 0) + m = MemoryManager.__new__(MemoryManager) + m.size = None + m.memory_reservations = reservations + m.get_cell_size = lambda: 4 + m.get_fixed_memory_size = lambda: 16 + m.profile_size(1) + assert m.size == expected From 782d52686eeb370f5061b5a903808d7a9e03c717 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 02:04:58 +0000 Subject: [PATCH 07/72] refactor: drop unused memory profiling reservations --- .../kv_cache_mem_manager/mem_manager.py | 43 +++---------------- .../models/deepseek_v4/test_memory_profile.py | 18 -------- 2 files changed, 5 insertions(+), 56 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index bc667b7d3a..9171566599 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -31,30 +31,13 @@ class MemoryManager: operator_class = NormalMemOperator - def __init__( - self, - size, - dtype, - head_num, - head_dim, - layer_num, - always_copy=False, - mem_fraction=0.9, - memory_reservations=None, - ): + def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False, mem_fraction=0.9): self.size = size self.head_num = head_num self.head_dim = head_dim self.layer_num = layer_num self.always_copy = always_copy self.dtype = dtype - # Named reservations are allocations made after KV profiling. They are - # deliberately outside get_fixed_memory_size(): model-specific exact KV - # geometry owns fixed bytes, while these values are deducted once from - # the profile budget only. - self.memory_reservations = dict(memory_reservations or {}) - if any(value < 0 for value in self.memory_reservations.values()): - raise ValueError(f"memory reservations must be non-negative: {self.memory_reservations}") # profile the max total token num if the size is None self.profile_size(mem_fraction) @@ -100,13 +83,6 @@ def get_pd_kv_move_buffer_size(self): ) return math.prod(shape) * torch._utils._element_size(self.dtype) - def get_fixed_memory_size(self): - return self.get_pd_kv_move_buffer_size() - - def get_profiled_size(self, available_memory_bytes): - """Select token capacity after fixed and post-profile reservations.""" - return int(available_memory_bytes / self.get_cell_size()) - def profile_size(self, mem_fraction): if self.size is not None: return @@ -115,25 +91,16 @@ def profile_size(self, mem_fraction): world_size = dist.get_world_size() available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) cell_size = self.get_cell_size() - fixed_memory_size = self.get_fixed_memory_size() - reservations = getattr(self, "memory_reservations", {}) - reserved_memory_size = sum(reservations.values()) - available_memory_bytes = available_memory * 1024 ** 3 - fixed_memory_size - reserved_memory_size - if available_memory_bytes <= 0: - raise RuntimeError( - f"{type(self).__name__} fixed buffers require {fixed_memory_size / 1024**3:.2f} GB, " - f"plus {reserved_memory_size / 1024**3:.2f} GB reservations, " - f"but only {available_memory:.2f} GB is available" - ) - self.size = self.get_profiled_size(available_memory_bytes) + pd_kv_move_buffer_size = self.get_pd_kv_move_buffer_size() + available_memory_bytes = available_memory * 1024 ** 3 - pd_kv_move_buffer_size + self.size = int(available_memory_bytes / cell_size) if world_size > 1: tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}") dist.all_reduce(tensor, op=dist.ReduceOp.MIN) self.size = tensor.item() logger.info( f"{str(available_memory)} GB space is available after load the model weight\n" - f"{str(fixed_memory_size / 1024 ** 2)} MB is reserved for fixed KV cache buffers\n" - f"{reservations} bytes are reserved for post-profile model buffers\n" + f"{str(pd_kv_move_buffer_size / 1024 ** 2)} MB is reserved for PD KV transfer buffer\n" f"{str(cell_size / 1024 ** 2)} MB is the size of one token kv cache\n" f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" ) diff --git a/unit_tests/models/deepseek_v4/test_memory_profile.py b/unit_tests/models/deepseek_v4/test_memory_profile.py index 4c25586831..f8abf6cc14 100644 --- a/unit_tests/models/deepseek_v4/test_memory_profile.py +++ b/unit_tests/models/deepseek_v4/test_memory_profile.py @@ -2,7 +2,6 @@ import pytest -from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager from lightllm.utils import profile_max_tokens @@ -45,20 +44,3 @@ def test_mtp_profile_exclusion_validation(monkeypatch, exclusion): with pytest.raises(ValueError, match="invalid MTP profile exclusion"): with profile_max_tokens.profile_mtp_weight_memory(model): pass - - -@pytest.mark.parametrize("reservations,expected", [({}, 252), ({"x": 20}, 247)]) -def test_memory_manager_profile_reservation_once(monkeypatch, reservations, expected): - monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.torch.cuda.empty_cache", lambda: None) - monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.dist.get_world_size", lambda: 1) - monkeypatch.setattr( - "lightllm.common.kv_cache_mem_manager.mem_manager.get_available_gpu_memory", lambda w: 1024 / 1024 ** 3 - ) - monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.get_total_gpu_memory", lambda: 0) - m = MemoryManager.__new__(MemoryManager) - m.size = None - m.memory_reservations = reservations - m.get_cell_size = lambda: 4 - m.get_fixed_memory_size = lambda: 16 - m.profile_size(1) - assert m.size == expected From 33d8db005c0a5d403560f6188bff36b32c1105cd Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 02:04:58 +0000 Subject: [PATCH 08/72] refactor: drop unused memory profiling hooks --- lightllm/utils/profile_max_tokens.py | 8 ---- .../models/deepseek_v4/test_memory_profile.py | 46 ------------------- 2 files changed, 54 deletions(-) delete mode 100644 unit_tests/models/deepseek_v4/test_memory_profile.py diff --git a/lightllm/utils/profile_max_tokens.py b/lightllm/utils/profile_max_tokens.py index 73f172d908..5ec2a9145e 100644 --- a/lightllm/utils/profile_max_tokens.py +++ b/lightllm/utils/profile_max_tokens.py @@ -74,14 +74,6 @@ def profile_mtp_weight_memory(model): weight_memory_before = torch.cuda.memory_allocated() yield target_weight_bytes = torch.cuda.memory_allocated() - weight_memory_before - # Draft construction can opt out of portions of the target's weight layout - # which it never instantiates (for example EPLB redundant rows). - excluded_weight_bytes = int(getattr(model, "get_mtp_profile_weight_exclusion", lambda: 0)()) - if not 0 <= excluded_weight_bytes <= target_weight_bytes: - raise ValueError( - f"invalid MTP profile exclusion {excluded_weight_bytes}; measured target weights={target_weight_bytes}" - ) - target_weight_bytes -= excluded_weight_bytes model.mem_fraction = get_mtp_adjusted_mem_fraction( mem_fraction=model.mem_fraction, target_weight_bytes=target_weight_bytes, diff --git a/unit_tests/models/deepseek_v4/test_memory_profile.py b/unit_tests/models/deepseek_v4/test_memory_profile.py deleted file mode 100644 index f8abf6cc14..0000000000 --- a/unit_tests/models/deepseek_v4/test_memory_profile.py +++ /dev/null @@ -1,46 +0,0 @@ -from types import SimpleNamespace - -import pytest - -from lightllm.utils import profile_max_tokens - - -@pytest.mark.parametrize("exclusion,expected", [(0, 1000), (200, 800), (None, 1000)]) -def test_mtp_profile_exclusion_adjustment(monkeypatch, exclusion, expected): - seen = [] - values = iter((100, 1100)) - monkeypatch.setattr(profile_max_tokens.torch.cuda, "memory_allocated", lambda: next(values)) - monkeypatch.setattr(profile_max_tokens, "get_mtp_weight_layer_num", lambda: 1) - monkeypatch.setattr( - profile_max_tokens, "get_mtp_adjusted_mem_fraction", lambda **kw: seen.append(kw["target_weight_bytes"]) or 0.5 - ) - attrs = dict( - max_total_token_num=None, - is_mtp_draft_model=False, - args=SimpleNamespace(mtp_mode="x"), - config={"n_layer": 1}, - mem_fraction=0.8, - ) - if exclusion is not None: - attrs["get_mtp_profile_weight_exclusion"] = lambda: exclusion - model = SimpleNamespace(**attrs) - with profile_max_tokens.profile_mtp_weight_memory(model): - pass - assert seen == [expected] - - -@pytest.mark.parametrize("exclusion", [-1, 1001]) -def test_mtp_profile_exclusion_validation(monkeypatch, exclusion): - values = iter((100, 1100)) - monkeypatch.setattr(profile_max_tokens.torch.cuda, "memory_allocated", lambda: next(values)) - model = SimpleNamespace( - max_total_token_num=None, - is_mtp_draft_model=False, - args=SimpleNamespace(mtp_mode="x"), - config={"n_layer": 1}, - mem_fraction=0.8, - get_mtp_profile_weight_exclusion=lambda: exclusion, - ) - with pytest.raises(ValueError, match="invalid MTP profile exclusion"): - with profile_max_tokens.profile_mtp_weight_memory(model): - pass From fd9dabeb2e18221df9cd8e08972595be44a6f9ca Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 02:31:09 +0000 Subject: [PATCH 09/72] refactor: derive EPLB enablement from redundant expert count --- .../fused_moe/fused_moe_weight.py | 4 +-- lightllm/distributed/communication_op.py | 6 +--- lightllm/server/api_cli.py | 11 ++----- lightllm/server/api_start.py | 14 ++++----- lightllm/server/core/objs/start_args_type.py | 3 +- .../model_infer/mode_backend/base_backend.py | 2 +- unit_tests/common/fused_moe/test_eplb.py | 12 ++++---- unit_tests/server/test_api_start_eplb.py | 30 +++++++++++++++++-- 8 files changed, 48 insertions(+), 34 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 0cd5a67522..131ab74e18 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 @@ -93,8 +93,8 @@ def _init_expert_parallel_state(self): self._initial_redundant_expert_ids = [] self._initial_redundant_expert_idx_to_local_idx = {} eplb = None - if args.enable_prefill_eplb and not is_eplb_model_init_disabled(): - num_redundant_experts_per_rank = args.eplb_num_redundant_experts_per_rank + num_redundant_experts_per_rank = args.eplb_num_redundant_experts_per_rank + if num_redundant_experts_per_rank > 0 and not is_eplb_model_init_disabled(): all_initial_ids = build_initial_redundant_expert_ids( self.n_routed_experts, self.global_world_size, diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index d18c4a780f..ee3411c4fc 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -201,11 +201,7 @@ def new_deepep_group( self.ll_num_tokens = prefill_num_max_dispatch_tokens_per_rank self.ll_decode_num_tokens = decode_num_max_dispatch_tokens_per_rank self.ll_hidden = hidden_size - total_redundant_experts = ( - get_env_start_args().eplb_num_redundant_experts_per_rank * global_world_size - if get_env_start_args().enable_prefill_eplb - else 0 - ) + total_redundant_experts = get_env_start_args().eplb_num_redundant_experts_per_rank * global_world_size self.ll_prefill_num_experts = n_routed_experts + total_redundant_experts # EPLB's redundant rows are a prefill-only physical layout; decode # always routes the logical expert space. diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index a2bb02879d..d615edcc72 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -771,17 +771,12 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="""Disable the prefill expert balance monitor enabled by default for EP-MoE.""", ) - parser.add_argument( - "--enable_prefill_eplb", - action="store_true", - help="""Enable online expert load balancing for prefill only.""", - ) parser.add_argument( "--eplb_num_redundant_experts_per_rank", type=int, - default=2, - help="""Number of redundant physical experts per EP rank for each MoE layer used by prefill EPLB. - The value must be greater than 0.""", + default=0, + help="""Number of redundant physical experts per EP rank for each MoE layer. + Set to 0 to disable EPLB.""", ) parser.add_argument( "--enable_fused_shared_experts", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 34ed299040..771dc9ed0c 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -159,14 +159,14 @@ def _launch_subprocesses(args: StartArgs): if args.enable_dp_prefill_balance: assert args.enable_tpsp_mix_mode and args.dp > 1, "need set --enable_tpsp_mix_mode firstly and --dp > 1" - if args.enable_prefill_eplb: - assert args.enable_ep_moe, "--enable_prefill_eplb requires --enable_ep_moe" - assert not args.enable_prefill_cudagraph, "--enable_prefill_eplb does not support --enable_prefill_cudagraph" + assert ( + args.eplb_num_redundant_experts_per_rank >= 0 + ), "--eplb_num_redundant_experts_per_rank must be greater than or equal to 0" + if args.eplb_num_redundant_experts_per_rank > 0: + assert args.enable_ep_moe, "EPLB requires --enable_ep_moe" + assert not args.enable_prefill_cudagraph, "EPLB does not support --enable_prefill_cudagraph" # EPLB updates expert weights in place, but SM100 Mega-MoE caches transformed weights by tensor data_ptr. - assert not is_sm100_gpu(), "--enable_prefill_eplb does not support SM100" - assert ( - args.eplb_num_redundant_experts_per_rank > 0 - ), "--eplb_num_redundant_experts_per_rank must be greater than 0" + assert not is_sm100_gpu(), "EPLB does not support SM100" if args.enable_ep_moe: allowed_ep_prefill_att_backends = {"auto", "fa3", "triton", "flashqla"} diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 0dd4881831..bec1d435c4 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -187,8 +187,7 @@ class StartArgs: ) enable_ep_moe: bool = field(default=False) disable_ep_balance_monitor: bool = field(default=False) - enable_prefill_eplb: bool = field(default=False) - eplb_num_redundant_experts_per_rank: int = field(default=2) + eplb_num_redundant_experts_per_rank: int = field(default=0) enable_fused_shared_experts: bool = field(default=False) mtp_mode: Optional[str] = field( default=None, diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 582a4470d6..081708083c 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -259,7 +259,7 @@ def init_model(self, kvargs): prof_name = f"lightllm-model_backend-node{self.node_rank}_dev{get_current_device_id()}" prof_mode = self.args.enable_profiling self.profiler = ProcessProfiler(mode=prof_mode, name=prof_name, use_multi_thread=True) if prof_mode else None - if self.args.enable_prefill_eplb: + if self.args.eplb_num_redundant_experts_per_rank > 0: from lightllm.server.router.model_infer.mode_backend.eplb_manager import EPLBManager self.model.eplb_manager = EPLBManager(self.model) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index fd656e07c3..f84ee1ef95 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -270,12 +270,12 @@ def __init__(self, layer_num, enable_ep_moe=True): assert manager_module._find_fused_moe_weights(model) == [alternate, aliased, first] -def test_eplb_redundant_experts_defaults_per_ep_rank(): +def test_eplb_redundant_experts_default_to_disabled(): parser = make_argument_parser() - assert parser.parse_args([]).eplb_num_redundant_experts_per_rank == 2 + assert parser.parse_args([]).eplb_num_redundant_experts_per_rank == 0 assert parser.parse_args(["--eplb_num_redundant_experts_per_rank", "3"]).eplb_num_redundant_experts_per_rank == 3 - assert StartArgs().eplb_num_redundant_experts_per_rank == 2 + assert StartArgs().eplb_num_redundant_experts_per_rank == 0 @pytest.mark.parametrize( @@ -1348,7 +1348,7 @@ def test_eplb_counter_capacity_covers_default_dense_interval(monkeypatch): args = type( "Args", (), - {"enable_prefill_eplb": True, "eplb_num_redundant_experts_per_rank": 2}, + {"eplb_num_redundant_experts_per_rank": 2}, )() weight = object.__new__(fused_weight_module.FusedMoeWeight) weight.n_routed_experts = 4 @@ -1399,7 +1399,7 @@ def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): args = type( "Args", (), - {"enable_prefill_eplb": False, "eplb_num_redundant_experts_per_rank": 2}, + {"eplb_num_redundant_experts_per_rank": 0}, )() weight = object.__new__(fused_weight_module.FusedMoeWeight) weight.n_routed_experts = 4 @@ -1422,7 +1422,7 @@ def test_disable_eplb_model_init_skips_eplb_state(monkeypatch): args = type( "Args", (), - {"enable_prefill_eplb": True, "eplb_num_redundant_experts_per_rank": 2}, + {"eplb_num_redundant_experts_per_rank": 2}, )() weight = object.__new__(fused_weight_module.FusedMoeWeight) weight.n_routed_experts = 4 diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py index 8064ab7952..8b988511ba 100644 --- a/unit_tests/server/test_api_start_eplb.py +++ b/unit_tests/server/test_api_start_eplb.py @@ -4,10 +4,34 @@ from lightllm.server.core.objs.start_args_type import StartArgs +@pytest.mark.parametrize( + ("redundant_experts", "message"), + [ + (-1, "--eplb_num_redundant_experts_per_rank must be greater than or equal to 0"), + (1, "EPLB requires --enable_ep_moe"), + ], +) +def test_eplb_redundant_expert_count_validation(monkeypatch, redundant_experts, message): + args = StartArgs( + eplb_num_redundant_experts_per_rank=redundant_experts, + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) + monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) + + with pytest.raises(AssertionError, match=message): + api_start._launch_subprocesses(args) + + def test_eplb_prefill_cudagraph_is_rejected_before_starting_subprocesses(monkeypatch): args = StartArgs( enable_ep_moe=True, - enable_prefill_eplb=True, + eplb_num_redundant_experts_per_rank=2, enable_prefill_cudagraph=True, disable_vision=True, disable_audio=True, @@ -24,7 +48,7 @@ def test_eplb_prefill_cudagraph_is_rejected_before_starting_subprocesses(monkeyp lambda *args, **kwargs: pytest.fail("subprocess startup must not be reached"), ) - with pytest.raises(AssertionError, match="--enable_prefill_eplb does not support --enable_prefill_cudagraph"): + with pytest.raises(AssertionError, match="EPLB does not support --enable_prefill_cudagraph"): api_start._launch_subprocesses(args) @@ -32,7 +56,7 @@ def test_eplb_mtp_combination_is_not_rejected_before_starting_subprocesses(monke args = StartArgs( model_dir="test-model", enable_ep_moe=True, - enable_prefill_eplb=True, + eplb_num_redundant_experts_per_rank=2, mtp_mode="vanilla_no_att", mtp_step=1, eos_id=0, From 7fa160610eeee27ce2d7a9c17a820c077fed0f32 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 02:34:52 +0000 Subject: [PATCH 10/72] test: require EP mode for EPLB redundancy --- unit_tests/server/test_api_start_eplb.py | 34 +++++++++++++++++------- 1 file changed, 24 insertions(+), 10 deletions(-) diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py index 8b988511ba..36c7fc48f6 100644 --- a/unit_tests/server/test_api_start_eplb.py +++ b/unit_tests/server/test_api_start_eplb.py @@ -4,16 +4,9 @@ from lightllm.server.core.objs.start_args_type import StartArgs -@pytest.mark.parametrize( - ("redundant_experts", "message"), - [ - (-1, "--eplb_num_redundant_experts_per_rank must be greater than or equal to 0"), - (1, "EPLB requires --enable_ep_moe"), - ], -) -def test_eplb_redundant_expert_count_validation(monkeypatch, redundant_experts, message): +def test_eplb_redundant_expert_count_must_not_be_negative(monkeypatch): args = StartArgs( - eplb_num_redundant_experts_per_rank=redundant_experts, + eplb_num_redundant_experts_per_rank=-1, disable_vision=True, disable_audio=True, disable_shm_warning=True, @@ -24,7 +17,28 @@ def test_eplb_redundant_expert_count_validation(monkeypatch, redundant_experts, monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) - with pytest.raises(AssertionError, match=message): + with pytest.raises( + AssertionError, + match="--eplb_num_redundant_experts_per_rank must be greater than or equal to 0", + ): + api_start._launch_subprocesses(args) + + +def test_eplb_redundant_experts_require_ep_moe(monkeypatch): + args = StartArgs( + enable_ep_moe=False, + eplb_num_redundant_experts_per_rank=1, + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) + monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) + + with pytest.raises(AssertionError, match="EPLB requires --enable_ep_moe"): api_start._launch_subprocesses(args) From 43db9afa3609acef6a6ce277d650129d0cd1000e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 02:42:25 +0000 Subject: [PATCH 11/72] refactor: unify DeepEP expert count management --- lightllm/distributed/communication_op.py | 30 +++++++++--------------- 1 file changed, 11 insertions(+), 19 deletions(-) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index ee3411c4fc..427cb80a8a 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -202,10 +202,7 @@ def new_deepep_group( self.ll_decode_num_tokens = decode_num_max_dispatch_tokens_per_rank self.ll_hidden = hidden_size total_redundant_experts = get_env_start_args().eplb_num_redundant_experts_per_rank * global_world_size - self.ll_prefill_num_experts = n_routed_experts + total_redundant_experts - # EPLB's redundant rows are a prefill-only physical layout; decode - # always routes the logical expert space. - self.ll_decode_num_experts = n_routed_experts + self.ll_num_experts = n_routed_experts + total_redundant_experts self.ep_buffer = deep_ep.ElasticBuffer( deepep_group, num_max_tokens_per_rank=self.ll_num_tokens, @@ -248,7 +245,7 @@ def new_deepep_group( self.ll_decode_num_tokens, self.ll_hidden, global_world_size, - self.ll_decode_num_experts, + self.ll_num_experts, ) microbatch_count = len(self.groups) min_prefill_reuse_buffer_bytes = _calculate_min_chunked_expanded_moe_reuse_buffer_bytes( @@ -269,7 +266,7 @@ def new_deepep_group( deepep_group, num_rdma_bytes=num_rdma_bytes, low_latency_mode=True, - num_qps_per_rank=(self.ll_decode_num_experts // global_world_size), + num_qps_per_rank=(self.ll_num_experts // global_world_size), ) if enable_mega_moe_buffer: @@ -282,7 +279,7 @@ def new_deepep_group( self.ep_mega_moe_buffer = deep_gemm.get_symm_buffer_for_mega_moe( deepep_group, - self.ll_decode_num_experts, + self.ll_num_experts, self.ll_num_tokens, num_experts_per_tok, self.ll_hidden, @@ -290,18 +287,16 @@ def new_deepep_group( ) logger.info( "Initialize DeepEP MoE buffers: low_latency=%s, mega_moe=%s, " - "ll_prefill_num_experts=%s, ll_decode_num_experts=%s, expert_quant_method_names=%s", + "ll_num_experts=%s, expert_quant_method_names=%s", enable_low_latency_buffer, enable_mega_moe_buffer, - self.ll_prefill_num_experts, - self.ll_decode_num_experts, + self.ll_num_experts, sorted(expert_quant_method_names), ) - theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_prefill_num_experts, num_experts_per_tok) - low_latency_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_decode_num_experts, num_experts_per_tok) - self._set_num_sms_for_deep_gemm(theoretical_sms, low_latency_sms) + 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) - def _set_num_sms_for_deep_gemm(self, deepep_sms: int, low_latency_sms: int): + def _set_num_sms_for_deep_gemm(self, deepep_sms: int): try: try: from deep_gemm.jit_kernels.utils import set_num_sms @@ -310,12 +305,9 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int, low_latency_sms: int): device_sms = get_device_sm_count() deepep_sms = max(0, min(deepep_sms, max(device_sms - 2, 0))) - low_latency_sms = max(0, min(low_latency_sms, max(device_sms - 2, 0))) self.ep_num_sms = deepep_sms if self.ep_low_latency_buffer is not None: - # This setting controls the legacy low-latency buffer; keep - # its SM reservation based on decode's logical expert count. - deep_ep.Buffer.set_num_sms(low_latency_sms - low_latency_sms % 2) + deep_ep.Buffer.set_num_sms(deepep_sms - deepep_sms % 2) set_num_sms(max(device_sms - deepep_sms, 2)) except BaseException as e: logger.warning(f"set num sms for deep_gemm failed: {e}") @@ -357,7 +349,7 @@ def clear_deepep_buffer(self): """ if self.ep_low_latency_buffer is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( - self.ll_decode_num_tokens, self.ll_hidden, self.ll_decode_num_experts + self.ll_decode_num_tokens, self.ll_hidden, self.ll_num_experts ) From 4ffd8122c14e07fe97786f3fc017a9736024f617 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 02:46:19 +0000 Subject: [PATCH 12/72] style: restore compact DeepEP size hint call --- lightllm/distributed/communication_op.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 427cb80a8a..a9ceef84f9 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -242,10 +242,7 @@ def new_deepep_group( # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 # 空闲的本地 RDMA storage 复用为分块 grouped GEMM 的临时 workspace。 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, + 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( From e9c3ba6e3d4702598cc6f197d5fd85416b54d068 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 02:56:40 +0000 Subject: [PATCH 13/72] refactor: unify EPLB routing across inference modes --- docs/CN/source/tutorial/api_server_args.rst | 26 ++++++++++++ docs/EN/source/tutorial/api_server_args.rst | 28 +++++++++++++ .../fused_moe/impl/deepgemm_impl.py | 38 ++--------------- .../triton_kernel/fused_moe/grouped_topk.py | 2 +- unit_tests/common/fused_moe/test_eplb.py | 41 +++++++++++-------- 5 files changed, 83 insertions(+), 52 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 53027e97a6..c45d8b9bac 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -705,6 +705,32 @@ PD 分离模式参数 使用 tgi 输入和输出格式 +专家并行与 EPLB 参数 +-------------------- + +.. option:: --enable_ep_moe + + 为支持的 MoE 模型启用专家并行。使用 EPLB 时必须开启此参数。 + +.. option:: --eplb_num_redundant_experts_per_rank + + 每个 MoE 层在每个 EP rank 上分配的冗余物理专家数量,默认值为 ``0``,表示关闭 EPLB。 + 设置为正数时启用 EPLB,并且必须同时设置 ``--enable_ep_moe``;负数会在启动阶段被拒绝。 + + 每个 rank 会额外分配指定数量的专家权重行。EPLB 将逻辑专家映射到主副本或冗余物理副本, + 统计路由负载,并可在线迁移冗余副本以改善专家负载均衡。增大此值可以提供更多布局选择, + 但也会占用更多 GPU 显存并增加专家迁移流量。 + + EPLB 当前不能与 ``--enable_prefill_cudagraph`` 同时使用,也不支持 SM100 GPU。 + 同一部署中的所有 rank 和节点必须使用相同的配置值。 + + 以下示例为每个 EP rank 配置两个冗余专家:: + + python -m lightllm.server.api_server \ + --model_dir /path/to/model \ + --enable_ep_moe \ + --eplb_num_redundant_experts_per_rank 2 + MTP 多预测参数 -------------- diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 0f82294c61..f24d16091f 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -721,6 +721,34 @@ Sampling and Generation Parameters Use tgi input and output format +Expert Parallelism and EPLB Parameters +-------------------------------------- + +.. option:: --enable_ep_moe + + Enable expert parallelism for supported MoE models. EPLB requires this option. + +.. option:: --eplb_num_redundant_experts_per_rank + + Number of redundant physical experts allocated on each EP rank for every MoE layer. The default is ``0``, + which disables EPLB. A positive value enables EPLB and must be used together with ``--enable_ep_moe``; + negative values are rejected during startup. + + Each rank allocates the configured number of additional expert weight rows. EPLB maps logical experts to + primary or redundant physical copies, records routing load, and can migrate redundant copies online to + improve expert load balance. Larger values provide more placement flexibility but consume more GPU memory + and increase expert migration traffic. + + EPLB currently cannot be combined with ``--enable_prefill_cudagraph`` and is not supported on SM100 GPUs. + Use the same value on every rank and node in one deployment. + + Example: enable EPLB with two redundant experts per EP rank:: + + python -m lightllm.server.api_server \ + --model_dir /path/to/model \ + --enable_ep_moe \ + --eplb_num_redundant_experts_per_rank 2 + MTP Multi-Prediction Parameters ------------------------------- 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 7ddd93fbb4..42c4dd4b3e 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 @@ -25,7 +25,6 @@ def __init__(self, *args, expert_parallel_state: ExpertParallelState, **kwargs): self.expert_parallel_state = expert_parallel_state self.eplb = expert_parallel_state.eplb self.ep_balance_counters = None - self._primary_weight_pack_cache = {} def _select_experts( self, @@ -43,10 +42,10 @@ def _select_experts( is_prefill: Optional[bool] = None, preserve_logical_ids: bool = False, ): - """选择 expert;EPLB prefill 统一由融合路径返回 physical ID。""" + """选择 expert;EPLB 统一由融合路径返回 physical ID。""" assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" eplb = self.eplb - if is_prefill is True and eplb is not None: + if eplb is not None: from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb group_score_topk_num = 2 if topk_group == 4 and num_expert_group == 8 and top_k == 8 else 1 @@ -97,20 +96,13 @@ def _fused_experts( router_logits: Optional[torch.Tensor] = None, is_prefill: Optional[bool] = None, ): - if is_prefill is False: - w13 = self._primary_weight_pack(w13) - w2 = self._primary_weight_pack(w2) - num_experts = self.n_routed_experts - else: - num_experts = self.expert_parallel_state.num_total_physical_experts - output = fused_experts( hidden_states=input_tensor, w13=w13, w2=w2, topk_weights=topk_weights, topk_idx=topk_ids.to(torch.long), - num_experts=num_experts, + num_experts=self.expert_parallel_state.num_total_physical_experts, quant_method=self.quant_method, is_prefill=is_prefill, previous_event=None, # for overlap @@ -150,8 +142,7 @@ def low_latency_dispatch( topk_idx=topk_idx, x=hidden_states, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, - # decode 与 EPLB 的物理冗余行刻意隔离:DeepEP 使用原始 logical expert ID。 - num_experts=self.n_routed_experts, + num_experts=self.expert_parallel_state.num_total_physical_experts, use_fp8=use_fp8_w8a8, async_finish=False, return_recv_hook=True, @@ -241,7 +232,6 @@ def masked_group_gemm( dtype: torch.dtype, expected_m: int, ): - w13, w2 = self._primary_weight_pack(w13), self._primary_weight_pack(w2) w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale return masked_group_gemm( @@ -339,23 +329,3 @@ def hook(): event.current_stream_wait() return combined_x, hook - - def _primary_weight_pack(self, weight_pack: WeightPack) -> WeightPack: - """返回所有 decode 路径使用的缓存本地主副本视图。""" - if self.eplb is None: - return weight_pack - cache = self._primary_weight_pack_cache - cache_key = id(weight_pack) - primary = cache.get(cache_key) - if primary is None: - num_primary_experts_per_rank = self.expert_parallel_state.num_primary_experts_per_rank - primary = WeightPack( - weight=weight_pack.weight[:num_primary_experts_per_rank], - weight_scale=( - weight_pack.weight_scale[:num_primary_experts_per_rank] - if weight_pack.weight_scale is not None - else None - ), - ) - cache[cache_key] = primary - return primary diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py index 7651282b2f..65883a5a64 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py @@ -456,7 +456,7 @@ def triton_grouped_topk_eplb( return_logical_ids: bool = False, group_score_used_topk_num: int = 2, ): - """Fused EPLB prefill top-k returning physical IDs and optional logical IDs.""" + """Fused EPLB top-k returning physical IDs and optional logical IDs.""" token_num, total_expert_num = gating_output.shape out_topk_weights = torch.empty((token_num, topk), dtype=torch.float32, device=gating_output.device) out_topk_ids = torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index f84ee1ef95..5218d6a109 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -124,7 +124,6 @@ def _validated_expert_parallel_state( def _set_expert_parallel_state(impl, state): impl.expert_parallel_state = state impl.eplb = state.eplb - impl._primary_weight_pack_cache = {} def _manual_runtime_rank_load(source_load, placement, node_world_size, alignment): @@ -1661,7 +1660,7 @@ def fail_prepare(_layer_plans): assert str(manager._evaluation_error) == "prepare failed" -def test_decode_dispatch_keeps_logical_ids_and_uses_logical_expert_count(monkeypatch): +def test_decode_dispatch_uses_physical_ids_and_total_expert_count(monkeypatch): class Buffer: def low_latency_dispatch(self, **kwargs): calls.append(kwargs) @@ -1673,7 +1672,7 @@ def low_latency_dispatch(self, **kwargs): _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) impl._select_experts = lambda **_kwargs: ( torch.ones((1, 2)), - torch.tensor([[0, 127]], dtype=torch.int32), + torch.tensor([[128, 143]], dtype=torch.int32), torch.tensor([[0, 127]], dtype=torch.int32), ) calls = [] @@ -1696,18 +1695,24 @@ def low_latency_dispatch(self, **kwargs): "softmax", ) - assert result[2].tolist() == [[0, 127]] - assert calls[0]["num_experts"] == 128 + assert result[2].tolist() == [[128, 143]] + assert calls[0]["num_experts"] == 144 -def test_decode_select_does_not_clone_or_map_eplb_topk_ids(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import topk_select +def test_decode_select_uses_eplb_mapping(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) impl.routed_scaling_factor = 1.0 _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) - topk_ids = torch.tensor([[3, 127]], dtype=torch.int32) - monkeypatch.setattr(topk_select, "select_experts", lambda **_kwargs: (torch.ones((1, 2)), topk_ids)) + physical_ids = torch.tensor([[130, 143]], dtype=torch.int32) + calls = [] + + def fused_topk(**kwargs): + calls.append(kwargs) + return torch.ones((1, 2)), physical_ids, None + + monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) _, selected, origin = impl._select_experts( torch.empty((1, 4)), torch.empty((1, 128)), @@ -1721,6 +1726,10 @@ def test_decode_select_does_not_clone_or_map_eplb_topk_ids(monkeypatch): is_prefill=False, ) + assert len(calls) == 1 + assert calls[0]["sample_index"] == 0 + assert not calls[0]["record_load"] + assert selected is physical_ids assert selected.data_ptr() == origin.data_ptr() @@ -1909,7 +1918,7 @@ def fused_topk(**kwargs): assert logical_ids.tolist() == [[3, 4]] -def test_decode_masked_group_gemm_uses_primary_rows_only_when_eplb_is_enabled( +def test_decode_masked_group_gemm_uses_all_physical_rows_when_eplb_is_enabled( monkeypatch, ): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) @@ -1931,11 +1940,11 @@ def masked(*args, **kwargs): )() assert impl.masked_group_gemm((torch.empty((1, 4)),), pack(), pack(), torch.empty(8), torch.float16, 1) == "out" - assert captured["w13"].shape[0] == captured["w2"].shape[0] == 8 - assert captured["w13_scale"].shape[0] == captured["w2_scale"].shape[0] == 8 + assert captured["w13"].shape[0] == captured["w2"].shape[0] == 10 + assert captured["w13_scale"].shape[0] == captured["w2_scale"].shape[0] == 10 -def test_decode_fused_experts_uses_cached_primary_weight_packs_and_logical_experts( +def test_decode_fused_experts_uses_full_weight_packs_and_physical_experts( monkeypatch, ): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) @@ -1977,10 +1986,8 @@ def fused(**kwargs): == "out" ) - assert [call["num_experts"] for call in captured] == [128, 128] - assert all(call["w13"].weight.shape[0] == call["w2"].weight.shape[0] == 8 for call in captured) - assert captured[0]["w13"] is captured[1]["w13"] - assert captured[0]["w2"] is captured[1]["w2"] + assert [call["num_experts"] for call in captured] == [160, 160] + assert all(call["w13"] is w13 and call["w2"] is w2 for call in captured) def test_transfer_plan_uses_existing_rows_and_prefers_local_node_replicas(): From b5ff3fde5b3c4fe13f8ee6d9905656c1220e96ca Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 03:04:44 +0000 Subject: [PATCH 14/72] refactor: remove inference phase from expert selection --- .../layer_weights/meta_weights/fused_moe/impl/base_impl.py | 2 -- .../meta_weights/fused_moe/impl/deepgemm_impl.py | 3 --- .../meta_weights/fused_moe/impl/triton_impl.py | 1 - unit_tests/common/fused_moe/test_eplb.py | 7 ++----- 4 files changed, 2 insertions(+), 11 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py index 35e872df10..a4758536d5 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py @@ -52,7 +52,6 @@ def __call__( scoring_func=scoring_func, per_expert_scale=per_expert_scale, shared_expert_gate=shared_expert_gate, - is_prefill=is_prefill, preserve_logical_ids=moe_capture_callback is not None, ) if moe_capture_callback is not None: @@ -81,7 +80,6 @@ def _select_experts( scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, shared_expert_gate: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, preserve_logical_ids: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: pass 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 42c4dd4b3e..56825ab1e8 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 @@ -39,7 +39,6 @@ def _select_experts( scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, shared_expert_gate: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, preserve_logical_ids: bool = False, ): """选择 expert;EPLB 统一由融合路径返回 physical ID。""" @@ -132,7 +131,6 @@ def low_latency_dispatch( topk_group=topk_group, num_expert_group=n_group, scoring_func=scoring_func, - is_prefill=False, ) topk_idx = topk_idx.to(torch.long) @@ -172,7 +170,6 @@ def select_experts_and_quant_input( topk_group=topk_group, num_expert_group=n_group, scoring_func=scoring_func, - is_prefill=True, ) qinput_tensor = quantize_fused_experts_input(hidden_states, w13, self.quant_method) return topk_weights, topk_idx.to(torch.long), qinput_tensor diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index c5f62ae946..e2bba71262 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -18,7 +18,6 @@ def _select_experts( scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, shared_expert_gate: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, preserve_logical_ids: bool = False, ): """Select experts and return topk weights and ids.""" diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 5218d6a109..f43f85b88c 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -164,7 +164,6 @@ def _select_experts( scoring_func, per_expert_scale=None, shared_expert_gate=None, - is_prefill=None, preserve_logical_ids=False, ): seen["select"] = {"preserve_logical_ids": preserve_logical_ids} @@ -1699,7 +1698,7 @@ def low_latency_dispatch(self, **kwargs): assert calls[0]["num_experts"] == 144 -def test_decode_select_uses_eplb_mapping(monkeypatch): +def test_select_uses_eplb_mapping(monkeypatch): from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) @@ -1723,7 +1722,6 @@ def fused_topk(**kwargs): 0, 0, "softmax", - is_prefill=False, ) assert len(calls) == 1 @@ -1884,7 +1882,7 @@ def test_deepgemm_constructor_configures_eplb(): assert impl.expert_parallel_state is state -def test_prefill_eplb_returns_requested_logical_ids(monkeypatch): +def test_eplb_select_returns_requested_logical_ids(monkeypatch): from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) @@ -1910,7 +1908,6 @@ def fused_topk(**kwargs): 0, 0, "softmax", - is_prefill=True, preserve_logical_ids=True, ) From 967daad950cc046ab9e96a93c05b6e3960854f3b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 05:14:05 +0000 Subject: [PATCH 15/72] refactor: separate MoE routing from execution layout --- .../meta_weights/fused_moe/impl/base_impl.py | 81 +++++- .../fused_moe/impl/deepgemm_impl.py | 68 ++--- .../fused_moe/impl/triton_impl.py | 13 +- .../triton_kernel/fused_moe/eplb_topk_ids.py | 90 ++++++ .../triton_kernel/fused_moe/grouped_topk.py | 251 ---------------- unit_tests/common/fused_moe/test_eplb.py | 272 ++++++------------ 6 files changed, 286 insertions(+), 489 deletions(-) create mode 100644 lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py index a4758536d5..fb4b1406f5 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py @@ -8,6 +8,32 @@ class FuseMoeBaseImpl(ABC): + """将逻辑专家路由与实际执行布局分离。 + + 融合 MoE 的调用流程如下:: + + _select_experts + -> topk_weights + logical_topk_ids + -> moe_capture_callback(logical_topk_ids) + -> _prepare_expert_execution + -> 追加 shared expert,或者 + -> 将 EPLB logical ID 映射为 physical expert ID + -> _fused_experts(topk_weights, execution_topk_ids) + + ``_select_experts`` 对所有实现都遵循同一套稳定接口:只在模型的逻辑专家 + 空间中选择专家,并返回原始 logical ID。该阶段不能应用任何与实际执行相关 + 的布局转换,例如追加 shared expert 行,或者映射到 EPLB 冗余 physical 行。 + + capture callback 在任何布局转换之前执行,因此采集到的路由元数据始终描述 + 模型的 logical expert。所有实现都必须将 logical ID 视为只读数据。 + ``_prepare_expert_execution`` 可以为实际执行布局分配新的 ID tensor,但不能 + 原地修改 logical ID tensor。 + + 因此,``_fused_experts`` 接收的是 execution ID:普通路径中仍是未经修改的 + logical ID;融合 shared expert 路径中是追加了 shared expert 行的 ID;EPLB + 路径中则是完成映射后的 physical ID。 + """ + def __init__( self, n_routed_experts: int, @@ -37,10 +63,12 @@ def __call__( # Callback to capture MoE topk expert ids (routed experts metadata). moe_capture_callback: Optional[Callable[[torch.Tensor], None]] = None, per_expert_scale: Optional[torch.Tensor] = None, - # Qwen3.5 uses this gate to control fused shared expert aggregation weights. + # Qwen3Next/Qwen3.5-MoE 在 TP 模式下将 shared expert 融合进 routed MoE。 + # 该参数是 shared_expert_gate(hidden_states) 产生的逐 token 门控 logit; + # 追加 shared expert 时使用 sigmoid(logit) 作为其聚合权重。 shared_expert_gate: Optional[torch.Tensor] = None, ) -> torch.Tensor: - topk_weights, topk_ids, origin_topk_ids = self._select_experts( + topk_weights, topk_ids = self._select_experts( input_tensor=input_tensor, router_logits=router_logits, correction_bias=correction_bias, @@ -51,11 +79,14 @@ def __call__( num_expert_group=num_expert_group, scoring_func=scoring_func, per_expert_scale=per_expert_scale, - shared_expert_gate=shared_expert_gate, - preserve_logical_ids=moe_capture_callback is not None, ) if moe_capture_callback is not None: - moe_capture_callback(origin_topk_ids) + moe_capture_callback(topk_ids) + topk_weights, topk_ids = self._prepare_expert_execution( + topk_weights=topk_weights, + topk_ids=topk_ids, + shared_expert_gate=shared_expert_gate, + ) return self._fused_experts( input_tensor=input_tensor, w13=w13, @@ -79,9 +110,38 @@ def _select_experts( num_expert_group: int, scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """在模型的逻辑专家空间中完成 top-k 路由选择。 + + 返回形状一致的 ``topk_weights`` 和 ``logical_topk_ids``。这里的 + ``topk_ids`` 必须始终表示模型配置中的原始 logical expert,不能追加 + shared expert,也不能映射到 EPLB physical expert。返回的 ID tensor 会 + 先交给 capture callback,后续实现必须将其视为只读数据。 + """ + pass + + @abstractmethod + def _prepare_expert_execution( + self, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, shared_expert_gate: Optional[torch.Tensor] = None, - preserve_logical_ids: bool = False, - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + ) -> Tuple[torch.Tensor, torch.Tensor]: + """将逻辑路由结果转换为 MoE kernel 实际需要的执行布局。 + + 该方法在 capture callback 之后执行。每个实现都必须明确处理自己的执行 + 布局:普通路径保持权重和 logical ID 不变,shared-expert 路径追加对应 + expert,EPLB 路径将 logical ID 修复为 physical ID。发生布局转换时应 + 返回新的 ID tensor,不能原地修改传入的 logical ``topk_ids``;如果专家 + 数量发生变化,``topk_weights`` 必须同步调整。 + + ``shared_expert_gate`` 当前由 Qwen3Next/Qwen3.5-MoE 使用,其内容是 + ``shared_expert_gate(hidden_states)`` 计算出的逐 token 门控 logit。在 TP + fused shared-expert 路径中,``sigmoid(logit)`` 会作为追加 shared expert + 的权重,使其输出按 token 动态参与 routed expert 输出的聚合;传入 ``None`` + 时,普通 fused shared expert 的追加权重为 1。EP 路径不使用该参数,而是 + 单独计算 shared expert,并在应用相同门控后与 routed MoE 输出相加。 + """ pass @abstractmethod @@ -95,4 +155,11 @@ def _fused_experts( router_logits: Optional[torch.Tensor] = None, is_prefill: Optional[bool] = None, ) -> torch.Tensor: + """根据准备完成的路由结果执行融合 MoE 计算。 + + 这里的 ``topk_ids`` 已处于实际执行所需的 ID 空间:普通路径为 logical + ID,shared expert 路径包含追加的专家 ID,EPLB 路径则为 physical ID。 + 实现只能读取路由 ID,不能原地修改其内容。``is_prefill`` 只用于选择底层 + 执行策略,不应再影响专家选择或 logical-to-physical 映射语义。 + """ pass 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 56825ab1e8..8fe8f4a17f 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 @@ -16,6 +16,7 @@ quantize_fused_experts_input, ) from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd +from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import eplb_repair_topk_ids from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType @@ -38,52 +39,43 @@ def _select_experts( num_expert_group: int, scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, + ): + """Select logical experts without applying the EPLB physical layout.""" + from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts + + topk_weights, topk_ids = select_experts( + hidden_states=input_tensor, + router_logits=router_logits, + correction_bias=correction_bias, + use_grouped_topk=use_grouped_topk, + top_k=top_k, + renormalize=renormalize, + topk_group=topk_group, + num_expert_group=num_expert_group, + scoring_func=scoring_func, + ) + if self.routed_scaling_factor != 1.0: + topk_weights.mul_(self.routed_scaling_factor) + return topk_weights, topk_ids + + def _prepare_expert_execution( + self, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, shared_expert_gate: Optional[torch.Tensor] = None, - preserve_logical_ids: bool = False, ): - """选择 expert;EPLB 统一由融合路径返回 physical ID。""" assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" eplb = self.eplb if eplb is not None: - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb - - group_score_topk_num = 2 if topk_group == 4 and num_expert_group == 8 and top_k == 8 else 1 - topk_weights, topk_ids, logical_topk_ids = triton_grouped_topk_eplb( - gating_output=router_logits, - correction_bias=correction_bias, - topk=top_k, - renormalize=renormalize, - num_expert_group=num_expert_group, - topk_group=topk_group, - scoring_func=scoring_func, + topk_ids = eplb_repair_topk_ids( + logical_topk_ids=topk_ids, logical_to_physical_map=eplb.logical_to_physical_map, logical_replica_count=eplb.logical_replica_count, expert_counter=eplb.route_counter, sample_index=eplb.next_sample_index(), record_load=eplb.recording, - use_grouped_topk=use_grouped_topk, - return_logical_ids=preserve_logical_ids, - group_score_used_topk_num=group_score_topk_num, ) - origin_topk_ids = logical_topk_ids if logical_topk_ids is not None else topk_ids - else: - from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts - - topk_weights, topk_ids = select_experts( - hidden_states=input_tensor, - router_logits=router_logits, - correction_bias=correction_bias, - use_grouped_topk=use_grouped_topk, - top_k=top_k, - renormalize=renormalize, - topk_group=topk_group, - num_expert_group=num_expert_group, - scoring_func=scoring_func, - ) - origin_topk_ids = topk_ids - if self.routed_scaling_factor != 1.0: - topk_weights.mul_(self.routed_scaling_factor) - return topk_weights, topk_ids, origin_topk_ids + return topk_weights, topk_ids def _fused_experts( self, @@ -121,7 +113,7 @@ def low_latency_dispatch( n_group: int, scoring_func: str, ): - topk_weights, topk_idx, _ = self._select_experts( + topk_weights, topk_idx = self._select_experts( input_tensor=hidden_states, router_logits=router_logits, correction_bias=e_score_correction_bias, @@ -132,6 +124,7 @@ def low_latency_dispatch( num_expert_group=n_group, scoring_func=scoring_func, ) + topk_weights, topk_idx = self._prepare_expert_execution(topk_weights, topk_idx) topk_idx = topk_idx.to(torch.long) num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_decode() @@ -160,7 +153,7 @@ def select_experts_and_quant_input( n_group: int, scoring_func: str, ): - topk_weights, topk_idx, _ = self._select_experts( + topk_weights, topk_idx = self._select_experts( input_tensor=hidden_states, router_logits=router_logits, correction_bias=e_score_correction_bias, @@ -171,6 +164,7 @@ def select_experts_and_quant_input( num_expert_group=n_group, scoring_func=scoring_func, ) + topk_weights, topk_idx = self._prepare_expert_execution(topk_weights, topk_idx) qinput_tensor = quantize_fused_experts_input(hidden_states, w13, self.quant_method) return topk_weights, topk_idx.to(torch.long), qinput_tensor diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index e2bba71262..a42f9c9f36 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -17,8 +17,6 @@ def _select_experts( num_expert_group: int, scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, - shared_expert_gate: Optional[torch.Tensor] = None, - preserve_logical_ids: bool = False, ): """Select experts and return topk weights and ids.""" from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts @@ -38,7 +36,14 @@ def _select_experts( topk_weights.mul_(self.routed_scaling_factor) if per_expert_scale is not None: topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype) - origin_topk_ids = topk_ids + return topk_weights, topk_ids + + def _prepare_expert_execution( + self, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + shared_expert_gate: Optional[torch.Tensor] = None, + ): if self.num_fused_shared_experts > 0: from lightllm.common.basemodel.triton_kernel.fused_moe.append_shared_expert_topk import ( append_fused_shared_experts, @@ -51,7 +56,7 @@ def _select_experts( num_fused_shared_experts=self.num_fused_shared_experts, shared_expert_gate=shared_expert_gate, ) - return topk_weights, topk_ids, origin_topk_ids + return topk_weights, topk_ids def _fused_experts( self, diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py new file mode 100644 index 0000000000..f0f491ec65 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py @@ -0,0 +1,90 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _replica_index(token_index, logical_id, replica_count): + token_hash = token_index.to(tl.uint32) * 2654435769 + expert_hash = logical_id.to(tl.uint32) * 2246822519 + return (token_hash + expert_hash) % replica_count.to(tl.uint32) + + +@triton.jit +def _eplb_repair_topk_ids_kernel( + logical_topk_ids_ptr, + physical_topk_ids_ptr, + topk_id_count, + topk, + logical_to_physical_map_ptr, + logical_replica_count_ptr, + expert_counter_ptr, + sample_index, + MAP_SLOTS: tl.constexpr, + COUNTER_NUM_EXPERTS: tl.constexpr, + RECORD_LOAD: tl.constexpr, + SINGLE_TOKEN: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < topk_id_count + logical_ids = tl.load(logical_topk_ids_ptr + offsets, mask=mask, other=0) + + if RECORD_LOAD: + tl.atomic_add( + expert_counter_ptr + sample_index * COUNTER_NUM_EXPERTS + logical_ids, + 1, + mask=mask, + sem="relaxed", + ) + + if SINGLE_TOKEN: + replica_indices = tl.zeros((BLOCK_SIZE,), tl.int32) + else: + replica_counts = tl.load(logical_replica_count_ptr + logical_ids, mask=mask, other=1) + token_indices = offsets // topk + replica_indices = _replica_index(token_indices, logical_ids, replica_counts) + + physical_ids = tl.load( + logical_to_physical_map_ptr + logical_ids * MAP_SLOTS + replica_indices, + mask=mask, + other=-1, + ) + tl.store(physical_topk_ids_ptr + offsets, physical_ids, mask=mask) + + +@torch.no_grad() +def eplb_repair_topk_ids( + logical_topk_ids: torch.Tensor, + logical_to_physical_map: torch.Tensor, + logical_replica_count: torch.Tensor, + expert_counter: torch.Tensor, + sample_index: int, + record_load: bool, +) -> torch.Tensor: + """Map logical top-k IDs to the current EPLB physical expert layout.""" + assert logical_topk_ids.is_contiguous() + assert logical_topk_ids.ndim == 2 + physical_topk_ids = torch.empty_like(logical_topk_ids) + if logical_topk_ids.numel() == 0: + return physical_topk_ids + + block_size = 512 + _eplb_repair_topk_ids_kernel[(triton.cdiv(logical_topk_ids.numel(), block_size),)]( + logical_topk_ids, + physical_topk_ids, + logical_topk_ids.numel(), + logical_topk_ids.shape[1], + logical_to_physical_map, + logical_replica_count, + expert_counter, + sample_index, + MAP_SLOTS=logical_to_physical_map.shape[1], + COUNTER_NUM_EXPERTS=expert_counter.shape[1], + RECORD_LOAD=record_load, + SINGLE_TOKEN=logical_topk_ids.shape[0] == 1, + BLOCK_SIZE=block_size, + num_warps=4, + num_stages=1, + ) + return physical_topk_ids diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py index 65883a5a64..fb0323cd4b 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py @@ -5,14 +5,6 @@ from triton.language.standard import _log2, sum, zeros_like -@triton.jit -def _eplb_replica_index(token_index, logical_id, replica_count): - """Choose a replica with independent phases for a token's top-k experts.""" - token_hash = token_index.to(tl.uint32) * 2654435769 - expert_hash = logical_id.to(tl.uint32) * 2246822519 - return (token_hash + expert_hash) % replica_count.to(tl.uint32) - - @triton.jit def _compare_and_swap(x, x_1, ids, flip, i: tl.core.constexpr, n_dims: tl.core.constexpr): n_outer: tl.core.constexpr = x.numel >> n_dims @@ -210,172 +202,6 @@ def grouped_topk_kernel( return -@triton.jit -def grouped_topk_eplb_kernel( - gating_output_ptr, - gating_output_stride_m, - gating_output_stride_n, - correction_bias_ptr, - out_topk_weights, - out_topk_weights_stride_m, - out_topk_weights_stride_n, - out_topk_ids, - out_topk_ids_stride_m, - out_topk_ids_stride_n, - out_logical_ids, - out_logical_ids_stride_m, - out_logical_ids_stride_n, - logical_to_physical_ptr, - logical_replica_count_ptr, - expert_counter_ptr, - sample_index, - group_num, - group_expert_num, - total_expert_num, - group_topk_num, - IS_SIGMOID: tl.constexpr, - USE_GROUPED_TOPK: tl.constexpr, - HAS_CORRECTION_BIAS: tl.constexpr, - RETURN_LOGICAL_IDS: tl.constexpr, - EXPERT_GROUP_NUM: tl.constexpr, - EXPERT_GROUP_SIZE: tl.constexpr, - TOPK_NUM: tl.constexpr, - TOPK_BLOCK_SIZE: tl.constexpr, - RENORMALIZE: tl.constexpr, - GROUP_SCORE_USED_TOPK_NUM: tl.constexpr, - COUNTER_NUM_EXPERTS: tl.constexpr, - MAP_SLOTS: tl.constexpr, - RECORD_LOAD: tl.constexpr, - SINGLE_TOKEN: tl.constexpr, -): - """Grouped top-k, EPLB accounting, and replica mapping without a global score workspace.""" - token_index = tl.program_id(axis=0) - offs_group = tl.arange(0, EXPERT_GROUP_NUM) - offs_group_v = tl.arange(0, EXPERT_GROUP_SIZE) - logical_ids = offs_group[:, None] * group_expert_num + offs_group_v[None, :] - valid_expert = ( - (offs_group < group_num)[:, None] - & (offs_group_v < group_expert_num)[None, :] - & (logical_ids < total_expert_num) - ) - hidden_states = tl.load( - gating_output_ptr + token_index * gating_output_stride_m + logical_ids * gating_output_stride_n, - mask=valid_expert, - other=-float("inf"), - ).to(tl.float32) - - if IS_SIGMOID: - old_scores = tl.sigmoid(hidden_states) - else: - group_max = tl.max(hidden_states, axis=1) - global_max = tl.max(group_max, axis=0) - numerators = tl.where(valid_expert, tl.exp(hidden_states - global_max), 0.0) - denominator = tl.sum(tl.sum(numerators, axis=1), axis=0) - old_scores = numerators / denominator - - if HAS_CORRECTION_BIAS: - correction_bias = tl.load(correction_bias_ptr + logical_ids, mask=valid_expert, other=0.0) - scores = tl.where(valid_expert, old_scores + correction_bias, -float("inf")) - else: - scores = tl.where(valid_expert, old_scores, -float("inf")) - - if USE_GROUPED_TOPK: - if GROUP_SCORE_USED_TOPK_NUM == 1: - group_value = tl.max(scores, axis=1) - elif GROUP_SCORE_USED_TOPK_NUM == 2: - first_score, first_index = tl.max(scores, axis=1, return_indices=True) - second_score = tl.max( - tl.where(offs_group_v[None, :] == first_index[:, None], -float("inf"), scores), - axis=1, - ) - group_value = first_score + second_score - else: - sorted_group_scores = tl.sort(scores, dim=1, descending=True) - group_value = tl.sum( - tl.where(offs_group_v[None, :] < GROUP_SCORE_USED_TOPK_NUM, sorted_group_scores, 0.0), - axis=1, - ) - - if EXPERT_GROUP_NUM > 1: - sorted_group_value = tl.sort(group_value, descending=True) - else: - sorted_group_value = group_value - group_topk_value = tl.sum(tl.where(offs_group == group_topk_num - 1, sorted_group_value, 0.0)) - candidate_scores = tl.where( - (group_value >= group_topk_value)[:, None] & valid_expert, - scores, - -float("inf"), - ) - else: - candidate_scores = tl.where(valid_expert, old_scores, -float("inf")) - - sort_block_size: tl.constexpr = EXPERT_GROUP_NUM * EXPERT_GROUP_SIZE - flat_offsets = tl.arange(0, sort_block_size) - candidate_scores = tl.reshape(candidate_scores, (sort_block_size,)) - topk_offsets = tl.arange(0, TOPK_BLOCK_SIZE) - selected_weights = tl.zeros((TOPK_BLOCK_SIZE,), tl.float32) - selected_logical_ids = tl.zeros((TOPK_BLOCK_SIZE,), tl.int32) - sum_scores = 0.0 - for topk_index in range(TOPK_NUM): - selected_offset = tl.argmax(candidate_scores, axis=0) - selected_group = selected_offset // EXPERT_GROUP_SIZE - selected_group_offset = selected_offset % EXPERT_GROUP_SIZE - selected_logical_id = selected_group * group_expert_num + selected_group_offset - selected_hidden_state = tl.load( - gating_output_ptr + token_index * gating_output_stride_m + selected_logical_id * gating_output_stride_n - ).to(tl.float32) - if IS_SIGMOID: - selected_weight = tl.sigmoid(selected_hidden_state) - else: - selected_weight = tl.exp(selected_hidden_state - global_max) / denominator - sum_scores += selected_weight - topk_lane = topk_offsets == topk_index - selected_weights = tl.where(topk_lane, selected_weight, selected_weights) - selected_logical_ids = tl.where(topk_lane, selected_logical_id, selected_logical_ids) - candidate_scores = tl.where(flat_offsets == selected_offset, -float("inf"), candidate_scores) - - topk_mask = topk_offsets < TOPK_NUM - if RECORD_LOAD: - tl.atomic_add( - expert_counter_ptr + sample_index * COUNTER_NUM_EXPERTS + selected_logical_ids, - 1, - mask=topk_mask, - sem="relaxed", - ) - if SINGLE_TOKEN: - replica_indices = tl.zeros((TOPK_BLOCK_SIZE,), tl.int32) - else: - replica_counts = tl.load( - logical_replica_count_ptr + selected_logical_ids, - mask=topk_mask, - other=1, - ) - replica_indices = _eplb_replica_index(token_index, selected_logical_ids, replica_counts) - selected_physical_ids = tl.load( - logical_to_physical_ptr + selected_logical_ids * MAP_SLOTS + replica_indices, - mask=topk_mask, - other=-1, - ) - if RENORMALIZE: - selected_weights /= sum_scores - tl.store( - out_topk_weights + token_index * out_topk_weights_stride_m + topk_offsets * out_topk_weights_stride_n, - selected_weights, - mask=topk_mask, - ) - tl.store( - out_topk_ids + token_index * out_topk_ids_stride_m + topk_offsets * out_topk_ids_stride_n, - selected_physical_ids, - mask=topk_mask, - ) - if RETURN_LOGICAL_IDS: - tl.store( - out_logical_ids + token_index * out_logical_ids_stride_m + topk_offsets * out_logical_ids_stride_n, - selected_logical_ids, - mask=topk_mask, - ) - - def triton_grouped_topk( hidden_states: torch.Tensor, gating_output: torch.Tensor, @@ -437,80 +263,3 @@ def triton_grouped_topk( num_stages=1, ) return out_topk_weights, out_topk_ids - - -def triton_grouped_topk_eplb( - gating_output: torch.Tensor, - correction_bias: torch.Tensor, - topk: int, - renormalize: bool, - num_expert_group: int, - topk_group: int, - scoring_func: str, - logical_to_physical_map: torch.Tensor, - logical_replica_count: torch.Tensor, - expert_counter: torch.Tensor, - sample_index: int, - record_load: bool, - use_grouped_topk: bool, - return_logical_ids: bool = False, - group_score_used_topk_num: int = 2, -): - """Fused EPLB top-k returning physical IDs and optional logical IDs.""" - token_num, total_expert_num = gating_output.shape - out_topk_weights = torch.empty((token_num, topk), dtype=torch.float32, device=gating_output.device) - out_topk_ids = torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) - out_logical_ids = ( - torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) if return_logical_ids else None - ) - if token_num == 0: - return out_topk_weights, out_topk_ids, out_logical_ids - if use_grouped_topk: - assert total_expert_num % num_expert_group == 0 - group_num = num_expert_group - group_expert_num = total_expert_num // num_expert_group - group_topk_num = topk_group - else: - group_num = 1 - group_expert_num = total_expert_num - group_topk_num = 1 - expert_group_num = triton.next_power_of_2(group_num) - expert_group_size = triton.next_power_of_2(group_expert_num) - sort_block_size = expert_group_num * expert_group_size - num_warps = min(max(1, sort_block_size // 256), 8) - grouped_topk_eplb_kernel[(token_num,)]( - gating_output, - *gating_output.stride(), - correction_bias, - out_topk_weights, - *out_topk_weights.stride(), - out_topk_ids, - *out_topk_ids.stride(), - out_logical_ids if out_logical_ids is not None else out_topk_ids, - *(out_logical_ids.stride() if out_logical_ids is not None else out_topk_ids.stride()), - logical_to_physical_map, - logical_replica_count, - expert_counter, - sample_index, - group_num=group_num, - group_expert_num=group_expert_num, - total_expert_num=total_expert_num, - group_topk_num=group_topk_num, - IS_SIGMOID=use_grouped_topk and scoring_func == "sigmoid", - USE_GROUPED_TOPK=use_grouped_topk, - HAS_CORRECTION_BIAS=use_grouped_topk and correction_bias is not None, - RETURN_LOGICAL_IDS=return_logical_ids, - EXPERT_GROUP_NUM=expert_group_num, - EXPERT_GROUP_SIZE=expert_group_size, - TOPK_NUM=topk, - TOPK_BLOCK_SIZE=triton.next_power_of_2(topk), - RENORMALIZE=renormalize, - GROUP_SCORE_USED_TOPK_NUM=group_score_used_topk_num, - COUNTER_NUM_EXPERTS=expert_counter.shape[1], - MAP_SLOTS=logical_to_physical_map.shape[1], - RECORD_LOAD=record_load, - SINGLE_TOKEN=token_num == 1, - num_warps=num_warps, - num_stages=1, - ) - return out_topk_weights, out_topk_ids, out_logical_ids diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index f43f85b88c..9b08f71857 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -163,11 +163,12 @@ def _select_experts( num_expert_group, scoring_func, per_expert_scale=None, - shared_expert_gate=None, - preserve_logical_ids=False, ): - seen["select"] = {"preserve_logical_ids": preserve_logical_ids} - return "weights", "physical_ids", "logical_ids" + return "weights", "logical_ids" + + def _prepare_expert_execution(self, topk_weights, topk_ids, shared_expert_gate=None): + seen["prepare"] = {"topk_ids": topk_ids} + return topk_weights, "physical_ids" def _fused_experts( self, @@ -200,7 +201,7 @@ def _fused_experts( ) assert result == "output" assert captured == ["logical_ids"] - assert seen["select"]["preserve_logical_ids"] + assert seen["prepare"]["topk_ids"] == "logical_ids" assert seen["fused"]["topk_ids"] == "physical_ids" @@ -1669,12 +1670,19 @@ def low_latency_dispatch(self, **kwargs): impl.quant_method = type("Quant", (), {"method_name": "fp8"})() impl.n_routed_experts = 128 _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + logical_ids = torch.tensor([[0, 127]], dtype=torch.int32) + physical_ids = torch.tensor([[128, 143]], dtype=torch.int32) impl._select_experts = lambda **_kwargs: ( torch.ones((1, 2)), - torch.tensor([[128, 143]], dtype=torch.int32), - torch.tensor([[0, 127]], dtype=torch.int32), + logical_ids, ) - calls = [] + calls, repairs = [], [] + + def repair(**kwargs): + repairs.append(kwargs) + return physical_ids + + monkeypatch.setattr(deepgemm_module, "eplb_repair_topk_ids", repair) monkeypatch.setattr( deepgemm_module, "get_deepep_num_max_dispatch_tokens_per_rank_decode", @@ -1695,24 +1703,25 @@ def low_latency_dispatch(self, **kwargs): ) assert result[2].tolist() == [[128, 143]] + assert repairs[0]["logical_topk_ids"] is logical_ids assert calls[0]["num_experts"] == 144 -def test_select_uses_eplb_mapping(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk +def test_select_returns_logical_ids_without_eplb_mapping(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import topk_select impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) impl.routed_scaling_factor = 1.0 _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) - physical_ids = torch.tensor([[130, 143]], dtype=torch.int32) + logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) calls = [] - def fused_topk(**kwargs): + def select(**kwargs): calls.append(kwargs) - return torch.ones((1, 2)), physical_ids, None + return torch.ones((1, 2)), logical_ids - monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) - _, selected, origin = impl._select_experts( + monkeypatch.setattr(topk_select, "select_experts", select) + _, selected = impl._select_experts( torch.empty((1, 4)), torch.empty((1, 128)), None, @@ -1725,27 +1734,24 @@ def fused_topk(**kwargs): ) assert len(calls) == 1 - assert calls[0]["sample_index"] == 0 - assert not calls[0]["record_load"] - assert selected is physical_ids - assert selected.data_ptr() == origin.data_ptr() - + assert selected is logical_ids -def test_eplb_prefill_uses_single_fused_path_for_global_topk(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk +def test_eplb_prefill_repairs_ids_after_selection(monkeypatch): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) impl.routed_scaling_factor = 1.0 impl.quant_method = object() _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) physical_ids = torch.tensor([[130, 131]], dtype=torch.long) calls = [] + impl._select_experts = lambda **_kwargs: (torch.ones((1, 2)), logical_ids) - def fused_topk(**kwargs): + def repair(**kwargs): calls.append(kwargs) - return torch.ones((1, 2)), physical_ids, None + return physical_ids - monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) + monkeypatch.setattr(deepgemm_module, "eplb_repair_topk_ids", repair) monkeypatch.setattr(deepgemm_module, "quantize_fused_experts_input", lambda *_args: "qinput") weights, topk_idx, qinput = impl.select_experts_and_quant_input( @@ -1765,13 +1771,11 @@ def fused_topk(**kwargs): assert topk_idx is physical_ids assert topk_idx.dtype is torch.long assert qinput == "qinput" - assert not calls[0]["use_grouped_topk"] - assert not calls[0]["return_logical_ids"] + assert calls[0]["logical_topk_ids"] is logical_ids + assert not calls[0]["record_load"] def test_eplb_prefill_dispatch_consumes_physical_ids_and_event(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk - class Buffer: def dispatch(self, _qinput, **kwargs): calls.append(kwargs) @@ -1793,14 +1797,16 @@ def dispatch(self, _qinput, **kwargs): ) _set_expert_parallel_state(impl, state) impl.ep_balance_counters = None - calls, fused_calls = [], [] + calls, repair_calls = [], [] + logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) physical_ids = torch.tensor([[130, 131]], dtype=torch.long) + impl._select_experts = lambda **_kwargs: (torch.ones((1, 2)), logical_ids) - def fused_topk(**kwargs): - fused_calls.append(kwargs) - return torch.ones((1, 2)), physical_ids, None + def repair(**kwargs): + repair_calls.append(kwargs) + return physical_ids - monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) + monkeypatch.setattr(deepgemm_module, "eplb_repair_topk_ids", repair) monkeypatch.setattr(deepgemm_module, "quantize_fused_experts_input", lambda *_args: "qinput") monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_buffer", Buffer()) monkeypatch.setattr( @@ -1831,9 +1837,10 @@ def fused_topk(**kwargs): ) assert topk_idx is physical_ids - assert len(fused_calls) == 1 - assert fused_calls[0]["sample_index"] == 0 - assert fused_calls[0]["record_load"] + assert len(repair_calls) == 1 + assert repair_calls[0]["logical_topk_ids"] is logical_ids + assert repair_calls[0]["sample_index"] == 0 + assert repair_calls[0]["record_load"] assert state.eplb.recorded_sample_count == 1 assert calls[0]["topk_idx"] is physical_ids assert calls[0]["topk_idx"].dtype is torch.long @@ -1882,37 +1889,25 @@ def test_deepgemm_constructor_configures_eplb(): assert impl.expert_parallel_state is state -def test_eplb_select_returns_requested_logical_ids(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk - +def test_eplb_prepare_repairs_logical_ids(monkeypatch): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - impl.routed_scaling_factor = 1.0 - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + state = _test_parallel_state(eplb=True, recording=True) + _set_expert_parallel_state(impl, state) + logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) + physical_ids = torch.tensor([[13, 14]], dtype=torch.int32) + calls = [] - def fused_topk(**kwargs): - assert kwargs["return_logical_ids"] - return ( - torch.ones((1, 2)), - torch.tensor([[13, 14]], dtype=torch.int32), - torch.tensor([[3, 4]], dtype=torch.int32), - ) + def repair(**kwargs): + calls.append(kwargs) + return physical_ids - monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) - _, physical_ids, logical_ids = impl._select_experts( - torch.empty((1, 4)), - torch.empty((1, 128)), - None, - 2, - False, - False, - 0, - 0, - "softmax", - preserve_logical_ids=True, - ) + monkeypatch.setattr(deepgemm_module, "eplb_repair_topk_ids", repair) + weights, selected = impl._prepare_expert_execution(torch.ones((1, 2)), logical_ids) - assert physical_ids.tolist() == [[13, 14]] - assert logical_ids.tolist() == [[3, 4]] + assert weights.tolist() == [[1.0, 1.0]] + assert selected is physical_ids + assert calls[0]["logical_topk_ids"] is logical_ids + assert calls[0]["record_load"] def test_decode_masked_group_gemm_uses_all_physical_rows_when_eplb_is_enabled( @@ -3492,21 +3487,13 @@ def new_group(*args, **kwargs): @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") @pytest.mark.parametrize("record_load", [False, True]) @pytest.mark.parametrize("tokens", [1, 32]) -@pytest.mark.parametrize("scoring_func", ["sigmoid", "softmax"]) -@pytest.mark.parametrize("renormalize", [False, True]) -def test_grouped_topk_eplb_matches_topk_mapping_and_counting(record_load, tokens, scoring_func, renormalize): - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import ( - triton_grouped_topk, - triton_grouped_topk_eplb, - ) - - torch.manual_seed(1234) - topk = 8 - experts = 256 - num_expert_group = 8 - gating_output = torch.randn((tokens, experts), dtype=torch.bfloat16, device="cuda") - correction_bias = torch.randn((experts,), dtype=torch.float32, device="cuda") - hidden_states = torch.empty((tokens, 1), dtype=torch.bfloat16, device="cuda") +def test_eplb_repair_topk_ids_maps_and_counts(record_load, tokens): + from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import eplb_repair_topk_ids + + topk = 4 + experts = 64 + logical_ids = (torch.arange(tokens * topk, dtype=torch.int32, device="cuda") % experts).view(tokens, topk) + original_logical_ids = logical_ids.clone() logical_to_physical = torch.stack( ( torch.arange(experts, dtype=torch.int32, device="cuda"), @@ -3519,20 +3506,9 @@ def test_grouped_topk_eplb_matches_topk_mapping_and_counting(record_load, tokens torch.full((experts,), 2, dtype=torch.int32, device="cuda"), torch.ones((experts,), dtype=torch.int32, device="cuda"), ) - expected_counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") - fused_counter = torch.zeros_like(expected_counter) - - expected_weights, logical_ids = triton_grouped_topk( - hidden_states, - gating_output, - correction_bias, - topk, - renormalize, - num_expert_group, - 4, - scoring_func, - 2, - ) + counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") + expected_counter = torch.zeros_like(counter) + if tokens == 1: replica_indices = torch.zeros_like(logical_ids) else: @@ -3541,129 +3517,45 @@ def test_grouped_topk_eplb_matches_topk_mapping_and_counting(record_load, tokens (((token_indices * 2654435769) & 0xFFFFFFFF) + ((logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF)) & 0xFFFFFFFF ) % logical_replica_count[logical_ids.to(torch.long)].to(torch.int64) - expected_ids = logical_to_physical[logical_ids.to(torch.long), replica_indices.to(torch.long)].to(torch.long) + expected_ids = logical_to_physical[logical_ids.to(torch.long), replica_indices.to(torch.long)] if record_load: expected_counter[1].scatter_add_( 0, logical_ids.reshape(-1).to(torch.long), torch.ones(logical_ids.numel(), dtype=torch.int64, device="cuda"), ) - fused_weights, fused_ids, fused_logical_ids = triton_grouped_topk_eplb( - gating_output, - correction_bias, - topk, - renormalize, - num_expert_group, - 4, - scoring_func, - logical_to_physical, - logical_replica_count, - fused_counter, - sample_index=1, - record_load=record_load, - use_grouped_topk=True, - group_score_used_topk_num=2, - ) - torch.cuda.synchronize() - - torch.testing.assert_close(fused_weights, expected_weights, rtol=1e-5, atol=1e-6) - assert torch.equal(fused_ids, expected_ids) - assert fused_logical_ids is None - assert torch.equal(fused_counter, expected_counter) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") -@pytest.mark.parametrize("record_load", [False, True]) -@pytest.mark.parametrize("tokens", [1, 32]) -def test_global_topk_eplb_supports_logical_ids_and_counting(record_load, tokens): - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb - torch.manual_seed(1234) - topk = 4 - experts = 64 - gating_output = torch.randn((tokens, experts), dtype=torch.float32, device="cuda") - logical_to_physical = torch.stack( - ( - torch.arange(experts, dtype=torch.int32, device="cuda"), - torch.arange(experts, dtype=torch.int32, device="cuda") + experts, - ), - dim=1, - ) - logical_replica_count = torch.where( - torch.arange(experts, device="cuda") % 3 == 0, - torch.full((experts,), 2, dtype=torch.int32, device="cuda"), - torch.ones((experts,), dtype=torch.int32, device="cuda"), - ) - expected_counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") - fused_counter = torch.zeros_like(expected_counter) - expected_weights, expected_logical_ids = torch.softmax(gating_output, dim=-1).topk(topk, dim=-1) - if tokens == 1: - replica_indices = torch.zeros_like(expected_logical_ids) - else: - token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) - replica_indices = ( - ( - ((token_indices * 2654435769) & 0xFFFFFFFF) - + ((expected_logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF) - ) - & 0xFFFFFFFF - ) % logical_replica_count[expected_logical_ids].to(torch.int64) - expected_ids = logical_to_physical[expected_logical_ids, replica_indices.to(torch.long)].to(torch.long) - if record_load: - expected_counter[1].scatter_add_( - 0, - expected_logical_ids.reshape(-1), - torch.ones(expected_logical_ids.numel(), dtype=torch.int64, device="cuda"), - ) - expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True) - - weights, physical_ids, logical_ids = triton_grouped_topk_eplb( - gating_output=gating_output, - correction_bias=torch.randn((experts,), dtype=torch.float32, device="cuda"), - topk=topk, - renormalize=True, - num_expert_group=8, - topk_group=4, - scoring_func="sigmoid", + physical_ids = eplb_repair_topk_ids( + logical_topk_ids=logical_ids, logical_to_physical_map=logical_to_physical, logical_replica_count=logical_replica_count, - expert_counter=fused_counter, + expert_counter=counter, sample_index=1, record_load=record_load, - use_grouped_topk=False, - return_logical_ids=True, ) torch.cuda.synchronize() - torch.testing.assert_close(weights, expected_weights, rtol=1e-5, atol=1e-6) + assert torch.equal(logical_ids, original_logical_ids) assert torch.equal(physical_ids, expected_ids) - assert torch.equal(logical_ids, expected_logical_ids) - assert torch.equal(fused_counter, expected_counter) + assert torch.equal(counter, expected_counter) @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") -def test_triton_grouped_topk_eplb_empty_tokens_skips_kernel(): - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb +def test_eplb_repair_topk_ids_empty_input_skips_kernel(): + from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import eplb_repair_topk_ids experts = 64 + logical_ids = torch.empty((0, 4), dtype=torch.int32, device="cuda") counter = torch.zeros((1, experts), dtype=torch.int64, device="cuda") - weights, physical_ids, logical_ids = triton_grouped_topk_eplb( - gating_output=torch.empty((0, experts), device="cuda"), - correction_bias=None, - topk=4, - renormalize=False, - num_expert_group=8, - topk_group=4, - scoring_func="softmax", + physical_ids = eplb_repair_topk_ids( + logical_topk_ids=logical_ids, logical_to_physical_map=torch.zeros((experts, 1), dtype=torch.int32, device="cuda"), logical_replica_count=torch.ones((experts,), dtype=torch.int32, device="cuda"), expert_counter=counter, sample_index=0, record_load=True, - use_grouped_topk=False, - return_logical_ids=True, ) - assert weights.shape == physical_ids.shape == logical_ids.shape == (0, 4) - assert physical_ids.dtype is logical_ids.dtype is torch.long + assert physical_ids.shape == (0, 4) + assert physical_ids.dtype is torch.int32 assert torch.equal(counter, torch.zeros_like(counter)) From 47bf4bbcff1a8cf0703819b9df7e94076ff3e7f0 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 05:34:10 +0000 Subject: [PATCH 16/72] refactor: simplify DeepGEMM MoE routing --- .../fused_moe/impl/deepgemm_impl.py | 18 ++++++++---------- unit_tests/common/fused_moe/test_eplb.py | 11 ++++++----- 2 files changed, 14 insertions(+), 15 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 8fe8f4a17f..453d0c1e47 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 @@ -24,7 +24,6 @@ class FuseMoeDeepGEMM(FuseMoeBaseImpl): def __init__(self, *args, expert_parallel_state: ExpertParallelState, **kwargs): super().__init__(*args, **kwargs) self.expert_parallel_state = expert_parallel_state - self.eplb = expert_parallel_state.eplb self.ep_balance_counters = None def _select_experts( @@ -56,6 +55,8 @@ def _select_experts( ) if self.routed_scaling_factor != 1.0: topk_weights.mul_(self.routed_scaling_factor) + if per_expert_scale is not None: + topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype) return topk_weights, topk_ids def _prepare_expert_execution( @@ -65,7 +66,7 @@ def _prepare_expert_execution( shared_expert_gate: Optional[torch.Tensor] = None, ): assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" - eplb = self.eplb + eplb = self.expert_parallel_state.eplb if eplb is not None: topk_ids = eplb_repair_topk_ids( logical_topk_ids=topk_ids, @@ -195,18 +196,15 @@ def dispatch( ) counters = self.ep_balance_counters - if counters is None: - - def hook(): - event.current_stream_wait() - - else: + route_load = compute_load = 0 + if counters is not None: # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. route_load = topk_idx.numel() compute_load = recv_x[0].shape[0] - def hook(): - event.current_stream_wait() + def hook(): + event.current_stream_wait() + if counters is not None: counters.accumulate( route_load=route_load, compute_load=compute_load, diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 9b08f71857..afa48c094c 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -123,7 +123,6 @@ def _validated_expert_parallel_state( def _set_expert_parallel_state(impl, state): impl.expert_parallel_state = state - impl.eplb = state.eplb def _manual_runtime_rank_load(source_load, placement, node_world_size, alignment): @@ -1707,21 +1706,21 @@ def repair(**kwargs): assert calls[0]["num_experts"] == 144 -def test_select_returns_logical_ids_without_eplb_mapping(monkeypatch): +def test_select_returns_logical_ids_and_applies_expert_scale(monkeypatch): from lightllm.common.basemodel.triton_kernel.fused_moe import topk_select impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - impl.routed_scaling_factor = 1.0 + impl.routed_scaling_factor = 2.0 _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) calls = [] def select(**kwargs): calls.append(kwargs) - return torch.ones((1, 2)), logical_ids + return torch.tensor([[0.5, 0.25]]), logical_ids monkeypatch.setattr(topk_select, "select_experts", select) - _, selected = impl._select_experts( + weights, selected = impl._select_experts( torch.empty((1, 4)), torch.empty((1, 128)), None, @@ -1731,9 +1730,11 @@ def select(**kwargs): 0, 0, "softmax", + per_expert_scale=torch.tensor([1.0, 1.0, 1.0, 3.0, 5.0]), ) assert len(calls) == 1 + assert weights.tolist() == [[3.0, 2.5]] assert selected is logical_ids From 487e44661eadbdcf1861ac5ba2b947764bfbea3c Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 10:40:34 +0000 Subject: [PATCH 17/72] refactor: simplify EPLB expert placement metadata --- .../meta_weights/fused_moe/eplb_placement.py | 474 +++++---- .../fused_moe/expert_parallel_state.py | 56 -- .../fused_moe/fused_moe_weight.py | 87 +- .../meta_weights/fused_moe/impl/__init__.py | 7 +- .../fused_moe/impl/deepgemm_impl.py | 95 +- .../triton_kernel/fused_moe/eplb_topk_ids.py | 114 ++- .../model_infer/mode_backend/base_backend.py | 6 +- .../model_infer/mode_backend/eplb_manager.py | 166 ++-- .../model_infer/mode_backend/eplb_transfer.py | 35 +- unit_tests/common/fused_moe/test_eplb.py | 926 +++++++++--------- .../fused_moe/test_eplb_transfer_gpu.py | 68 +- 11 files changed, 998 insertions(+), 1036 deletions(-) delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py index 1ea8eb1643..b43e011b8b 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -1,85 +1,196 @@ -from dataclasses import dataclass -from functools import lru_cache from typing import Dict, Tuple import torch -def build_initial_redundant_expert_ids( +def build_initial_local_expert_ids( num_logical_experts: int, num_ranks: int, num_redundant_experts_per_rank: int, -) -> torch.Tensor: - """Build a deterministic initial placement without local duplicates.""" +) -> list[list[int]]: + """构建每个 rank 初始持有的完整 logical expert ID 列表。 + + 每个 rank 先持有连续划分得到的主专家,再用本 rank 最后一个主专家 + 填充额外物理槽。这里仅负责生成 Python 列表;调用方如果要参与 tensor + 运算,需要自行转换为 ``torch.Tensor``。 + + 例如 ``num_logical_experts=8``、``num_ranks=4``、每个 rank 有 2 个 + 额外槽时,每个 rank 分到 2 个主专家,结果为: + + ``[[0, 1, 1, 1], [2, 3, 3, 3], [4, 5, 5, 5], [6, 7, 7, 7]]`` + + 其中每行前两个值是主专家,后两个值是等待首次 EPLB 调整的占位副本。 + """ assert num_logical_experts % num_ranks == 0 num_experts_per_rank = num_logical_experts // num_ranks - assert 0 < num_redundant_experts_per_rank <= num_logical_experts - num_experts_per_rank + assert num_redundant_experts_per_rank >= 0 + + local_expert_ids_by_rank = [] + for rank in range(num_ranks): + first_expert_id = rank * num_experts_per_rank + local_expert_ids = list(range(first_expert_id, first_expert_id + num_experts_per_rank)) + # 额外物理槽先复制本 rank 最后一个主专家。首次 EPLB 规划完成后, + # transfer 会把这些占位行替换成实际需要的跨 rank 冗余专家。 + local_expert_ids.extend([local_expert_ids[-1]] * num_redundant_experts_per_rank) + local_expert_ids_by_rank.append(local_expert_ids) - # 初始化结果确定,不依赖随机数。 - # 每个 rank 不会复制自己原本拥有的 expert。 - # 同一个 rank 的冗余槽位不会重复。 - # 最后一个 rank 通过取模自然回绕。 - rank_offsets = torch.arange(1, num_ranks + 1, dtype=torch.int64)[:, None] * num_experts_per_rank - expert_offsets = torch.arange(num_redundant_experts_per_rank, dtype=torch.int64) - return (rank_offsets + expert_offsets) % num_logical_experts + return local_expert_ids_by_rank def build_logical_to_physical_map( - redundant_expert_ids: torch.Tensor, # 冗余布局,shape 为 [num_ranks, num_redundant_experts_per_rank]。 - num_logical_experts: int, # 逻辑 expert 的总数。 - source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 - node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 -) -> Tuple[ - torch.Tensor, torch.Tensor -]: # logical_to_physical [num_logical_experts, num_ranks], replica_counts [num_logical_experts] - """构建单层逻辑 expert 到物理副本的映射""" - - logical_to_physical, replica_counts = build_logical_to_physical_maps_for_layers( - redundant_expert_ids.unsqueeze(0), + rank_to_logic_expert_ids: list[list[int]], + num_logical_experts: int, + current_rank: int, +) -> list[list[int]]: + """使用普通 CPU list 构建单层 logical 到 physical expert 的路由表。 + + ``rank_to_logic_expert_ids`` 的 shape 为 + ``[num_ranks, num_physical_experts_per_rank]``,每行包含该 rank 的全部 + 主专家和冗余专家。 + + 返回值的 shape 为 ``[num_logical_experts, 2 + routing_slots]``。每一行 + 对应一个 logical expert:第 0 项是有效副本数,第 1 项 + 标记 ``current_rank`` 是否持有本地副本,第 2 项起是 physical expert ID。 + 如果本 rank 持有副本,该副本固定放在第一个路由槽;有效副本之后未使用 + 的固定宽度 padding 槽位填充为 ``-1``。 + + 本函数只负责 CPU 元数据计算。调用方需要设备 Tensor 时,应在函数外 + 显式执行 ``torch.tensor(...)``。 + """ + # 阶段 1:校验输入布局,并根据每个 rank 的物理槽位数计算冗余容量。 + num_ranks = len(rank_to_logic_expert_ids) + assert num_ranks > 0 + assert num_logical_experts % num_ranks == 0 + num_physical_experts_per_rank = len(rank_to_logic_expert_ids[0]) + assert all(len(rank_expert_ids) == num_physical_experts_per_rank for rank_expert_ids in rank_to_logic_expert_ids) + num_primary_experts_per_rank = num_logical_experts // num_ranks + num_redundant_experts_per_rank = num_physical_experts_per_rank - num_primary_experts_per_rank + assert num_redundant_experts_per_rank >= 0 + # 阶段 2:计算固定路由槽宽度。最坏情况下,所有 rank 的全部冗余槽都 + # 指向同一个 logical expert;再加上该 expert 固有的一个主副本,就是 + # 任意 logical expert 可能拥有的最大物理副本数。 + num_routing_slots = 1 + num_ranks * num_redundant_experts_per_rank + assert 0 <= current_rank < num_ranks + + # 阶段 3:把“物理槽 -> logical expert”的完整布局反转为 + # “logical expert -> 全部物理槽”,得到每个专家的候选副本列表。 + physical_ids_by_logical_expert = _collect_physical_ids_by_logical_expert( + rank_to_logic_expert_ids, num_logical_experts, - source_rank=source_rank, - node_world_size=node_world_size, ) - return logical_to_physical.squeeze(0), replica_counts.squeeze(0) + + # 阶段 4:对每个候选列表做稳定排序。本 rank 的 physical ID 排在前面, + # 因而后续只需查看第一个候选,就能判断和选择本地副本。 + _sort_physical_ids_by_locality( + physical_ids_by_logical_expert, + current_rank, + num_physical_experts_per_rank, + ) + local_physical_id_start = current_rank * num_physical_experts_per_rank + local_physical_id_end = local_physical_id_start + num_physical_experts_per_rank + + # 阶段 5:逐个 logical expert 打包固定宽度的路由行。实际副本不足固定 + # 宽度时,剩余槽位使用 -1 padding;kernel 只会索引有效副本范围。 + logical_to_physical_map = [] + for physical_expert_ids in physical_ids_by_logical_expert: + has_local_replica = local_physical_id_start <= physical_expert_ids[0] < local_physical_id_end + logical_to_physical_map.append( + _build_routing_row( + physical_expert_ids=physical_expert_ids, + has_local_replica=has_local_replica, + num_routing_slots=num_routing_slots, + ) + ) + return logical_to_physical_map + + +def _collect_physical_ids_by_logical_expert( + rank_to_logic_expert_ids: list[list[int]], + num_logical_experts: int, +) -> list[list[int]]: + """将完整物理布局反转为每个 logical expert 对应的物理槽位。 + + 例如输入 ``[[0, 1, 1], [2, 3, 0]]``,先按 rank 顺序拼成 + ``[0, 1, 1, 2, 3, 0]``。该列表的下标就是 physical expert ID,值就是 + logical expert ID,因此最终返回 ``[[0, 5], [1, 2], [3], [4]]``。 + """ + logical_expert_ids_by_physical_id = [ + logical_expert_id for rank_expert_ids in rank_to_logic_expert_ids for logical_expert_id in rank_expert_ids + ] + physical_ids_by_logical_expert = [[] for _ in range(num_logical_experts)] + for physical_expert_id, logical_expert_id in enumerate(logical_expert_ids_by_physical_id): + assert 0 <= logical_expert_id < num_logical_experts + physical_ids_by_logical_expert[logical_expert_id].append(physical_expert_id) + + return physical_ids_by_logical_expert + + +def _sort_physical_ids_by_locality( + physical_ids_by_logical_expert: list[list[int]], + current_rank: int, + num_physical_experts_per_rank: int, +) -> None: + """按照 physical ID 是否属于当前 rank,对每个副本列表稳定排序。 + + 本地 physical ID 的排序键为 0,其他 physical ID 的排序键为 1。因此 + 当前 rank 持有的副本会移动到列表前面,同时本地副本之间、远端副本 + 之间的原始顺序保持不变。当前 rank 没有副本的列表顺序不会发生变化。 + """ + local_physical_id_start = current_rank * num_physical_experts_per_rank + local_physical_id_end = local_physical_id_start + num_physical_experts_per_rank + + for physical_expert_ids in physical_ids_by_logical_expert: + # list.sort 是稳定排序:排序键相同时,physical ID 的原始顺序不变。 + physical_expert_ids.sort( + key=lambda physical_expert_id: ( + 0 if local_physical_id_start <= physical_expert_id < local_physical_id_end else 1 + ) + ) + + +def _build_routing_row( + physical_expert_ids: list[int], + has_local_replica: bool, + num_routing_slots: int, +) -> list[int]: + """将一个 logical expert 的候选 physical IDs 打包为固定宽度路由行。 + + ``physical_expert_ids`` 已由调用方完成本地优先的稳定排序,所以本函数 + 不再依赖 ``current_rank``。列表长度就是该 logical expert 的有效物理 + 副本数,无需额外传入容易失配的副本数量。 + """ + # 阶段 1:候选列表包含一个主副本及全部冗余副本,其长度就是有效副本数。 + num_valid_replicas = len(physical_expert_ids) + assert 0 < num_valid_replicas <= num_routing_slots + + # 阶段 2:有效槽位直接保存稳定排序后的候选;固定宽度中未使用的尾部 + # 槽位统一填充 -1。kernel 的副本索引严格小于 num_valid_replicas, + # 因而不会读取 padding。 + num_padding_slots = num_routing_slots - num_valid_replicas + routing_slots = physical_expert_ids + [-1] * num_padding_slots + + # 阶段 3:第 0 列保存 kernel 参与 hash 的有效副本数;第 1 列标记是否 + # 存在本地副本;后续列保存按本地优先顺序排列的 physical IDs 和 -1 padding。 + return [num_valid_replicas, int(has_local_replica), *routing_slots] def build_logical_to_physical_maps_for_layers( - redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] - num_logical_experts: int, # 逻辑 expert 的总数。 - source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 - node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 -) -> Tuple[ - torch.Tensor, # logical_to_physical, shape [num_layers, num_logical_experts, num_ranks] - torch.Tensor, # replica_counts, shape [num_layers, num_logical_experts] -]: - """Build stable CPU int32 maps for the supplied layers without modifying the input.""" - if redundant_expert_ids_by_layer.ndim != 3: - raise ValueError("redundant_expert_ids_by_layer must be [layers, ranks, num_redundant_experts_per_rank]") - num_ranks, num_redundant_experts_per_rank = redundant_expert_ids_by_layer.shape[1:] - assert num_logical_experts % num_ranks == 0 - layout = _get_physical_expert_layout(num_logical_experts, num_ranks, num_redundant_experts_per_rank) - logical_to_physical, replica_counts = _build_global_replica_maps_for_layers(redundant_expert_ids_by_layer, layout) - if source_rank is None: - return logical_to_physical, replica_counts - - assert node_world_size is not None - replica_positions = torch.arange(num_ranks, dtype=torch.int64) - compact_maps_by_layer, selected_counts_by_layer = _select_source_node_replicas( - logical_to_physical, - replica_counts, - source_rank=source_rank, - node_world_size=node_world_size, - num_physical_experts_per_rank=layout.num_physical_experts_per_rank, - replica_positions=replica_positions, - ) - return ( - _rotate_selected_replicas( - compact_maps_by_layer, - selected_counts_by_layer, - source_rank=source_rank, - replica_positions=replica_positions, - ), - selected_counts_by_layer, - ) + rank_to_logic_expert_ids_by_layer: list[list[list[int]]], + num_logical_experts: int, + current_rank: int, +) -> list[list[list[int]]]: + """逐层构建 CPU list 路由表;设备 Tensor 由调用方在边界处创建。 + + 输入 shape 为 ``[num_layers, num_ranks, num_physical_experts_per_rank]``, + 输出 shape 为 ``[num_layers, num_logical_experts, 2 + routing_slots]``。 + """ + return [ + build_logical_to_physical_map( + rank_to_logic_expert_ids, + num_logical_experts, + current_rank=current_rank, + ) + for rank_to_logic_expert_ids in rank_to_logic_expert_ids_by_layer + ] def select_improving_placements( @@ -89,14 +200,13 @@ def select_improving_placements( *, rebalance_gain_threshold: float, expert_alignment: int | None = None, - node_world_size: int | None = None, ) -> Tuple[torch.Tensor, torch.Tensor, Dict[str, float | int], torch.Tensor, torch.Tensor]: """Select better layers and return current/final rank loads without re-estimation.""" if not 0.0 <= rebalance_gain_threshold <= 1.0: raise ValueError("rebalance_gain_threshold must be between 0.0 and 1.0") assert current_placement.shape == candidate_placement.shape - current_rank_load = _estimate_rank_load(expert_load, current_placement, expert_alignment, node_world_size) - candidate_rank_load = _estimate_rank_load(expert_load, candidate_placement, expert_alignment, node_world_size) + current_rank_load = _estimate_rank_load(expert_load, current_placement, expert_alignment) + candidate_rank_load = _estimate_rank_load(expert_load, candidate_placement, expert_alignment) current_critical = current_rank_load.max(dim=-1).values.sum(dim=0) candidate_critical = candidate_rank_load.max(dim=-1).values.sum(dim=0) # Each changed layer must reduce its own critical load. All selected @@ -137,11 +247,10 @@ def plan_redundant_experts( num_ranks: int, num_redundant_experts_per_rank: int, expert_alignment: int | None = None, - node_world_size: int | None = None, current_placement: torch.Tensor | None = None, stickiness: float = 0.0, ) -> torch.Tensor: - """Plan replicas from [samples, layers, source_nodes, experts] loads. + """Plan replicas from [samples, layers, current_ranks, experts] loads. With ``current_placement`` and positive ``stickiness``, a candidate that keeps an expert on its current rank receives a bonus of @@ -154,8 +263,9 @@ def plan_redundant_experts( """ if expert_alignment is not None: assert expert_alignment > 0 - node_world_size = _resolve_node_world_size(expert_load, num_ranks, node_world_size) - num_samples, num_layers, num_nodes, num_logical_experts = expert_load.shape + assert expert_load.ndim == 4 + _, num_layers, num_load_ranks, num_logical_experts = expert_load.shape + assert num_load_ranks in (1, num_ranks) assert num_logical_experts % num_ranks == 0 assert num_redundant_experts_per_rank > 0 num_experts_per_rank = num_logical_experts // num_ranks @@ -178,7 +288,7 @@ def plan_redundant_experts( stickiness_scale = None locations = _expert_locations(placement, num_logical_experts) - expert_rank = _expert_rank_load_all(load, locations, num_nodes, node_world_size, expert_alignment) + expert_rank = _expert_rank_load_all(load, locations, expert_alignment) rank_load = expert_rank.sum(dim=2) remaining_slots = torch.full((num_layers, num_ranks), num_redundant_experts_per_rank, dtype=torch.int64) layer_indices = torch.arange(num_layers, dtype=torch.int64) @@ -204,9 +314,7 @@ def plan_redundant_experts( candidate_locations = locations.clone() candidate_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] = True - candidate_expert_rank = _expert_rank_load_all( - load, candidate_locations, num_nodes, node_world_size, expert_alignment - ) + candidate_expert_rank = _expert_rank_load_all(load, candidate_locations, expert_alignment) candidate_rank_load = rank_load[:, :, None, :] - expert_rank + candidate_expert_rank critical = candidate_rank_load.max(dim=3).values.sum(dim=0) critical.masked_fill_(~legal, torch.inf) @@ -234,235 +342,89 @@ def plan_redundant_experts( return placement -@dataclass(frozen=True, eq=False) -class _PhysicalExpertLayout: - """进程内按拓扑复用的只读物理 expert 布局;其中 Tensor 不得原地修改。""" - - num_logical_experts: int - num_ranks: int - num_physical_experts_per_rank: int - primary_physical_ids: torch.Tensor - redundant_physical_ids: torch.Tensor - - -@lru_cache(maxsize=8) -def _get_physical_expert_layout( - num_logical_experts: int, - num_ranks: int, - num_redundant_experts_per_rank: int, -) -> _PhysicalExpertLayout: - """返回按静态拓扑缓存的只读 CPU 物理 expert ID。""" - num_experts_per_rank = num_logical_experts // num_ranks - num_physical_experts_per_rank = num_experts_per_rank + num_redundant_experts_per_rank - expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) - primary_physical_ids = ( - (expert_ids // num_experts_per_rank) * num_physical_experts_per_rank + expert_ids % num_experts_per_rank - ).to(torch.int32) - ranks = torch.arange(num_ranks, dtype=torch.int64).repeat_interleave(num_redundant_experts_per_rank) - slots = torch.arange(num_redundant_experts_per_rank, dtype=torch.int64).repeat(num_ranks) - redundant_physical_ids = (ranks * num_physical_experts_per_rank + num_experts_per_rank + slots).to(torch.int32) - return _PhysicalExpertLayout( - num_logical_experts=num_logical_experts, - num_ranks=num_ranks, - num_physical_experts_per_rank=num_physical_experts_per_rank, - primary_physical_ids=primary_physical_ids, - redundant_physical_ids=redundant_physical_ids, - ) - - -def _build_global_replica_maps_for_layers( - redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] - layout: _PhysicalExpertLayout, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Build stable global maps with primary copies first and unused slots set to ``-1``.""" - - num_layers = redundant_expert_ids_by_layer.shape[0] - num_logical_experts = layout.num_logical_experts - max_replicas = layout.num_ranks - redundant_ids = redundant_expert_ids_by_layer.to(dtype=torch.int64, device="cpu") - logical_to_physical = torch.full((num_layers, num_logical_experts, max_replicas), -1, dtype=torch.int32) - logical_to_physical[:, :, 0] = layout.primary_physical_ids - replica_counts = torch.ones((num_layers, num_logical_experts), dtype=torch.int32) - - flat_redundant_ids = redundant_ids.reshape(num_layers, -1) - if not flat_redundant_ids.numel(): - return logical_to_physical, replica_counts - - # 稳定排序保留 rank-major、slot-major 的历史顺序;第 0 列固定为主副本。 - sort_order = torch.argsort(flat_redundant_ids, dim=1, stable=True) - sorted_redundant_ids = flat_redundant_ids.gather(1, sort_order) - flat_positions = torch.arange(flat_redundant_ids.shape[1], dtype=torch.int64).unsqueeze(0) - group_starts = torch.where( - torch.cat( - ( - torch.ones((num_layers, 1), dtype=torch.bool), - sorted_redundant_ids[:, 1:] != sorted_redundant_ids[:, :-1], - ), - dim=1, - ), - flat_positions, - 0, - ) - replica_indices = flat_positions - torch.cummax(group_starts, dim=1).values + 1 - redundant_counts = torch.zeros((num_layers, num_logical_experts), dtype=torch.int32) - redundant_counts.scatter_add_( - 1, - flat_redundant_ids, - torch.ones_like(flat_redundant_ids, dtype=torch.int32), - ) - assert int(redundant_counts.max().item()) < max_replicas, "an expert can have at most one replica per rank" - replica_counts += redundant_counts - - layer_indices = torch.arange(num_layers, dtype=torch.int64).view(-1, 1).expand_as(sort_order) - redundant_physical_ids = layout.redundant_physical_ids.unsqueeze(0).expand_as(sort_order).gather(1, sort_order) - logical_to_physical[layer_indices, sorted_redundant_ids, replica_indices] = redundant_physical_ids - return logical_to_physical, replica_counts - - -def _select_source_node_replicas( - logical_to_physical: torch.Tensor, - replica_counts: torch.Tensor, - *, - source_rank: int, - node_world_size: int, - num_physical_experts_per_rank: int, - replica_positions: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Keep source-node replicas when available, otherwise fall back to all stable candidates.""" - num_layers, num_logical_experts, _max_replicas = logical_to_physical.shape - source_node = source_rank // node_world_size - output_positions = replica_positions.view(1, 1, -1) - valid = output_positions < replica_counts.unsqueeze(-1) - local = valid & ( - torch.div( - logical_to_physical, - num_physical_experts_per_rank * node_world_size, - rounding_mode="floor", - ) - == source_node - ) - selected = torch.where(local.any(dim=2, keepdim=True), local, valid) - selected_counts_by_layer = selected.sum(dim=2, dtype=torch.int32) - - compact_maps_by_layer = torch.full_like(logical_to_physical, -1) - selected_positions = selected.cumsum(dim=2) - 1 - layers = torch.arange(num_layers, dtype=torch.int64).view(-1, 1, 1).expand_as(selected) - experts = torch.arange(num_logical_experts, dtype=torch.int64).view(1, -1, 1).expand_as(selected) - compact_maps_by_layer[layers[selected], experts[selected], selected_positions[selected]] = logical_to_physical[ - selected - ] - return compact_maps_by_layer, selected_counts_by_layer - - -def _rotate_selected_replicas( - compact_maps_by_layer: torch.Tensor, - selected_count_by_layer: torch.Tensor, - *, - source_rank: int, - replica_positions: torch.Tensor, -) -> torch.Tensor: - """根据 source_rank 循环调整每个逻辑 expert 的候选副本顺序,让不同源 rank 优先使用不同副本,同时保持候选副本集合和副本数量不变""" - output_positions = replica_positions.view(1, 1, -1) - selected_count64_by_layer = selected_count_by_layer.to(torch.int64).unsqueeze(-1) - source_positions_by_layer = (output_positions + source_rank) % selected_count64_by_layer - maps_by_layer = compact_maps_by_layer.gather(2, source_positions_by_layer) - maps_by_layer.masked_fill_(output_positions >= selected_count64_by_layer, -1) - return maps_by_layer - - def _estimate_rank_load( expert_load: torch.Tensor, - redundant_expert_ids: torch.Tensor, + rank_to_logic_expert_ids: torch.Tensor, expert_alignment: int | None = None, - node_world_size: int | None = None, ) -> torch.Tensor: - """Estimate [samples, layers, ranks] load from source-node-local routing. + """Estimate [samples, layers, ranks] load from current-rank-local routing. - Source loads remain separate until assigned to physical replicas, then + Per-rank loads remain separate until assigned to physical replicas, then combine before the per-expert alignment used by DeepEP. """ - node_world_size = _resolve_node_world_size(expert_load, redundant_expert_ids.shape[1], node_world_size) - num_samples, num_layers, num_nodes, num_logical_experts = expert_load.shape - assert redundant_expert_ids.ndim == 3 and redundant_expert_ids.shape[0] == num_layers - num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape[1:] + assert expert_load.ndim == 4 + _, num_layers, num_load_ranks, num_logical_experts = expert_load.shape + assert rank_to_logic_expert_ids.ndim == 3 and rank_to_logic_expert_ids.shape[0] == num_layers + num_ranks = rank_to_logic_expert_ids.shape[1] + assert num_load_ranks in (1, num_ranks) assert num_logical_experts % num_ranks == 0 if expert_alignment is not None: assert expert_alignment > 0 rank_load = _expert_rank_load_all( expert_load, - _expert_locations(redundant_expert_ids, num_logical_experts), - num_nodes, - node_world_size, + _expert_locations(rank_to_logic_expert_ids, num_logical_experts), expert_alignment, ).sum(dim=2) return rank_load -def _resolve_node_world_size(expert_load: torch.Tensor, num_ranks: int, node_world_size: int | None) -> int: - """Validate production [samples, layers, source_nodes, experts] planner loads.""" - assert expert_load.ndim == 4 - num_nodes = expert_load.shape[2] - if node_world_size is None: - assert num_ranks % num_nodes == 0 - node_world_size = num_ranks // num_nodes - assert 0 < node_world_size <= num_ranks and num_ranks % node_world_size == 0 - assert num_nodes == num_ranks // node_world_size - return node_world_size - - -def _expert_locations(redundant_expert_ids: torch.Tensor, num_logical_experts: int) -> torch.Tensor: +def _expert_locations(rank_to_logic_expert_ids: torch.Tensor, num_logical_experts: int) -> torch.Tensor: """Return ``[layer, logical expert, rank]`` physical-copy occupancy.""" - num_layers, num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape + num_layers, num_ranks, num_redundant_experts_per_rank = rank_to_logic_expert_ids.shape assert num_logical_experts % num_ranks == 0 num_experts_per_rank = num_logical_experts // num_ranks locations = torch.zeros( (num_layers, num_logical_experts, num_ranks), dtype=torch.bool, - device=redundant_expert_ids.device, + device=rank_to_logic_expert_ids.device, ) expert_ids = torch.arange(num_logical_experts, device=locations.device) owners = expert_ids // num_experts_per_rank locations[:, expert_ids, owners] = True layers = torch.arange(num_layers, device=locations.device)[:, None] ranks = torch.arange(num_ranks, device=locations.device).repeat_interleave(num_redundant_experts_per_rank)[None, :] - redundant_ids = redundant_expert_ids.reshape(num_layers, -1) - valid = redundant_ids >= 0 + flat_logic_expert_ids = rank_to_logic_expert_ids.reshape(num_layers, -1) + valid = flat_logic_expert_ids >= 0 if torch.any(valid): - expanded_layers = layers.expand_as(redundant_ids) - expanded_ranks = ranks.expand_as(redundant_ids) + expanded_layers = layers.expand_as(flat_logic_expert_ids) + expanded_ranks = ranks.expand_as(flat_logic_expert_ids) locations[ expanded_layers[valid], - redundant_ids[valid], + flat_logic_expert_ids[valid], expanded_ranks[valid], ] = True return locations -def _source_route(slots: torch.Tensor, num_nodes: int, node_world_size: int) -> torch.Tensor: - """Route each source node to its local copies, or all copies as fallback.""" +def _current_rank_route(slots: torch.Tensor, num_load_ranks: int) -> torch.Tensor: + """当前 rank 有本地副本时只选本地,否则在所有副本之间均分。 + + ``num_load_ranks == 1`` 表示离线调用方只提供了聚合负载,此时无法判断 + 当前 rank,直接在全部副本之间均分。线上采样始终传入逐 rank 负载。 + """ num_ranks = slots.shape[-1] - assert num_ranks % node_world_size == 0 and num_nodes == num_ranks // node_world_size - rank_nodes = torch.arange(num_ranks, device=slots.device) // node_world_size - source_nodes = torch.arange(num_nodes, device=slots.device) - copies = slots.unsqueeze(-3).expand(*slots.shape[:-2], num_nodes, *slots.shape[-2:]) - rank_node_shape = (1,) * slots.ndim + (num_ranks,) - source_node_shape = (1,) * (slots.ndim - 2) + (num_nodes, 1, 1) - local = copies & (rank_nodes.reshape(rank_node_shape) == source_nodes.reshape(source_node_shape)) + copies = slots.unsqueeze(-3).expand(*slots.shape[:-2], num_load_ranks, *slots.shape[-2:]) + if num_load_ranks == 1: + return copies.to(torch.float64) / copies.sum(dim=-1, keepdim=True) + + assert num_load_ranks == num_ranks + ranks = torch.arange(num_ranks, device=slots.device) + destination_rank_shape = (1,) * slots.ndim + (num_ranks,) + current_rank_shape = (1,) * (slots.ndim - 2) + (num_ranks, 1, 1) + local = copies & (ranks.reshape(destination_rank_shape) == ranks.reshape(current_rank_shape)) selected = torch.where(local.any(dim=-1, keepdim=True), local, copies) return selected.to(torch.float64) / selected.sum(dim=-1, keepdim=True) def _expert_rank_load_all( - source_load: torch.Tensor, + expert_load: torch.Tensor, locations: torch.Tensor, - num_nodes: int, - node_world_size: int, expert_alignment: int | None, ) -> torch.Tensor: """Return aligned ``[samples, layers, expert, rank]`` contributions.""" - route = _source_route(locations, num_nodes, node_world_size) - physical_load = torch.einsum("slne,lner->sler", source_load.to(torch.float64), route) + route = _current_rank_route(locations, expert_load.shape[2]) + physical_load = torch.einsum("slqe,lqer->sler", expert_load.to(torch.float64), route) if expert_alignment is not None: physical_load = torch.ceil(physical_load / expert_alignment) * expert_alignment return physical_load diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py deleted file mode 100644 index f66d31a62b..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py +++ /dev/null @@ -1,56 +0,0 @@ -from contextlib import contextmanager -from contextvars import ContextVar -from dataclasses import dataclass -from typing import Iterator, Optional - -import torch - - -_eplb_model_init_disabled: ContextVar[bool] = ContextVar("eplb_model_init_disabled", default=False) - - -def is_eplb_model_init_disabled() -> bool: - return _eplb_model_init_disabled.get() - - -@contextmanager -def disable_eplb_model_init() -> Iterator[None]: - token = _eplb_model_init_disabled.set(True) - try: - yield - finally: - _eplb_model_init_disabled.reset(token) - - -@dataclass -class EPLBState: - num_redundant_experts_per_rank: int - initial_redundant_expert_ids_by_rank: torch.Tensor - logical_to_physical_map: torch.Tensor - logical_replica_count: torch.Tensor - route_counter: torch.Tensor - recording: bool = False - recorded_sample_count: int = 0 - - def next_sample_index(self) -> int: - if not self.recording: - return 0 - sample_index = self.recorded_sample_count % self.route_counter.shape[0] - self.recorded_sample_count += 1 - return sample_index - - -@dataclass(frozen=True) -class ExpertParallelState: - num_logical_experts: int - world_size: int - eplb: Optional[EPLBState] = None - - @property - def num_primary_experts_per_rank(self) -> int: - return self.num_logical_experts // self.world_size - - @property - def num_total_physical_experts(self) -> int: - num_redundant_experts_per_rank = 0 if self.eplb is None else self.eplb.num_redundant_experts_per_rank - return self.num_logical_experts + self.world_size * num_redundant_experts_per_rank 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 131ab74e18..7b34c60d0a 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 @@ -9,19 +9,10 @@ SliceMixinTpl, ) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import create_fuse_moe_impl -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - EPLBState, - ExpertParallelState, - is_eplb_model_init_disabled, -) from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_initial_redundant_expert_ids, - build_logical_to_physical_map, -) from lightllm.common.quantization.quantize_method import QuantizationMethod -from lightllm.utils.envs_utils import get_env_start_args, get_prefill_eplb_step_interval -from lightllm.utils.dist_utils import get_global_world_size, get_global_rank, get_node_world_size +from lightllm.utils.envs_utils import 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 logger = init_logger(__name__) @@ -65,15 +56,14 @@ def __init__( self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self._init_config(network_config) - self._init_expert_parallel_state() - self._init_weight_partition() self.fuse_moe_impl = create_fuse_moe_impl( n_routed_experts=self.n_routed_experts, num_fused_shared_experts=self.num_fused_shared_experts, routed_scaling_factor=self.routed_scaling_factor, quant_method=self.quant_method, - expert_parallel_state=self.expert_parallel_state, + enable_ep_moe=self.enable_ep_moe, ) + self._init_weight_partition() self.lock = threading.Lock() self._create_weight() @@ -86,49 +76,6 @@ def _init_config(self, network_config: Dict[str, Any]): self.routed_scaling_factor = network_config.get("routed_scaling_factor", 1.0) self.scoring_func = network_config.get("scoring_func", "softmax") - def _init_expert_parallel_state(self): - args = get_env_start_args() - self.expert_parallel_state: Optional[ExpertParallelState] = None - # Initial placement metadata is used only while loading checkpoint rows. - self._initial_redundant_expert_ids = [] - self._initial_redundant_expert_idx_to_local_idx = {} - eplb = None - num_redundant_experts_per_rank = args.eplb_num_redundant_experts_per_rank - if num_redundant_experts_per_rank > 0 and not is_eplb_model_init_disabled(): - all_initial_ids = build_initial_redundant_expert_ids( - self.n_routed_experts, - self.global_world_size, - num_redundant_experts_per_rank, - ) - self._initial_redundant_expert_ids = all_initial_ids[self.global_rank_].tolist() - logical_to_physical, logical_replica_count = build_logical_to_physical_map( - all_initial_ids, - self.n_routed_experts, - source_rank=self.global_rank_, - node_world_size=get_node_world_size(), - ) - # route_counter 每次 prefill dispatch 记录一行。初始阶段连续采样 - # step_interval 个 manager step,兼顾micro batch overlap的两次 dispatch,因此容量设为 - # 2 * step_interval。稳定阶段复用该环形缓冲区,但只把当前短采样窗口内 - # 实际记录的最近行传给 planner,不复制整个缓冲区。 - eplb = EPLBState( - num_redundant_experts_per_rank=num_redundant_experts_per_rank, - initial_redundant_expert_ids_by_rank=all_initial_ids, - logical_to_physical_map=logical_to_physical.cuda(), - logical_replica_count=logical_replica_count.cuda(), - route_counter=torch.zeros( - (2 * get_prefill_eplb_step_interval(), self.n_routed_experts), - dtype=torch.int64, - device="cuda", - ), - ) - if self.enable_ep_moe: - self.expert_parallel_state = ExpertParallelState( - num_logical_experts=self.n_routed_experts, - world_size=self.global_world_size, - eplb=eplb, - ) - def _init_weight_partition(self): if self.enable_ep_moe: self.tp_rank_ = 0 @@ -143,26 +90,14 @@ def _init_weight_partition(self): self.split_inter_size = self.moe_intermediate_size // self.tp_world_size_ if self.enable_ep_moe: assert self.num_fused_shared_experts == 0, "num_fused_shared_experts must be 0 when enable_ep_moe" - eplb = self.expert_parallel_state.eplb - num_redundant_experts_per_rank = 0 if eplb is None else eplb.num_redundant_experts_per_rank + self.local_expert_ids = self.fuse_moe_impl.local_logics_expert_ids_list logger.debug( f"global_rank {self.global_rank_} layerindex {self.layer_num_} " - f"initial_redundant_expert_ids: {self._initial_redundant_expert_ids}" + f"local_logics_expert_ids_list: {self.local_expert_ids}" ) - num_primary_experts_per_rank = self.expert_parallel_state.num_primary_experts_per_rank - self.local_n_routed_experts = num_primary_experts_per_rank + num_redundant_experts_per_rank - start_expert_id = self.global_rank_ * num_primary_experts_per_rank - self.expert_idx_to_local_idx = { - expert_idx: expert_idx - start_expert_id - for expert_idx in range(start_expert_id, start_expert_id + num_primary_experts_per_rank) - } - self._initial_redundant_expert_idx_to_local_idx = { - redundant_expert_idx: num_primary_experts_per_rank + i - for (i, redundant_expert_idx) in enumerate(self._initial_redundant_expert_ids) - } + self.local_n_routed_experts = len(self.local_expert_ids) else: self.local_expert_ids = list(range(self.n_routed_experts + self.num_fused_shared_experts)) - self.expert_idx_to_local_idx = {expert_idx: i for (i, expert_idx) in enumerate(self.local_expert_ids)} def experts( self, @@ -319,9 +254,7 @@ def load_hf_weights(self, weights): # Load bias self._load_e_score_correction_bias(weights) self._load_per_expert_scale(weights) - self._load_weight(self.expert_idx_to_local_idx, weights) - if self._initial_redundant_expert_idx_to_local_idx: - self._load_weight(self._initial_redundant_expert_idx_to_local_idx, weights) + self._load_weight(self.local_expert_ids, weights) def verify_load(self): weight_load_ok = all(all(_weight_pack.load_ok) for _weight_pack in self.w1_list + self.w2_list + self.w3_list) @@ -390,8 +323,8 @@ def _get_expert_weight_list(self, weight_pack: WeightPack): weight_list.append(expert_weight) return weight_list - def _load_weight(self, expert_idx_to_local_idx: Dict[int, int], weights: Dict[str, torch.Tensor]): - for expert_idx, local_expert_idx in expert_idx_to_local_idx.items(): + def _load_weight(self, local_expert_ids: List[int], weights: Dict[str, torch.Tensor]): + for local_expert_idx, expert_idx in enumerate(local_expert_ids): with self.lock: self._load_expert(expert_idx, local_expert_idx, weights) self._load_expert_scale( diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index acbc31d261..20522b5c37 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -2,7 +2,6 @@ from .triton_impl import FuseMoeTriton from .marlin_impl import FuseMoeMarlin from .deepgemm_impl import FuseMoeDeepGEMM -from ..expert_parallel_state import ExpertParallelState def create_fuse_moe_impl( @@ -11,9 +10,9 @@ def create_fuse_moe_impl( num_fused_shared_experts: int, routed_scaling_factor: float, quant_method: QuantizationMethod, - expert_parallel_state: ExpertParallelState | None = None, + enable_ep_moe: bool = False, ): - if expert_parallel_state is not None: + if enable_ep_moe: impl_cls = FuseMoeDeepGEMM elif quant_method.method_name == "awq_marlin": impl_cls = FuseMoeMarlin @@ -25,6 +24,4 @@ def create_fuse_moe_impl( routed_scaling_factor=routed_scaling_factor, quant_method=quant_method, ) - if expert_parallel_state is not None: - kwargs["expert_parallel_state"] = expert_parallel_state return impl_cls(**kwargs) 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 453d0c1e47..f6ef8e8e0a 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 @@ -1,13 +1,21 @@ import torch from typing import Optional, Tuple, Any from .base_impl import FuseMoeBaseImpl -from ..expert_parallel_state import ExpertParallelState +from ..eplb_placement import ( + build_initial_local_expert_ids, + build_logical_to_physical_map, +) from lightllm.distributed import dist_group_manager from lightllm.common.quantization.quantize_method import WeightPack from lightllm.utils.envs_utils import ( + get_env_start_args, get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, ) +from lightllm.utils.dist_utils import ( + get_global_rank, + get_global_world_size, +) from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( fused_experts, get_ep_num_sms, @@ -15,17 +23,57 @@ 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.basemodel.triton_kernel.fused_moe.eplb_topk_ids import eplb_repair_topk_ids +from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import ( + silu_and_mul_fwd, +) +from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( + eplb_repair_topk_ids, +) from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType class FuseMoeDeepGEMM(FuseMoeBaseImpl): - def __init__(self, *args, expert_parallel_state: ExpertParallelState, **kwargs): + def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.expert_parallel_state = expert_parallel_state + self._init_eplb_runtime() self.ep_balance_counters = None + def _init_eplb_runtime(self): + world_size = get_global_world_size() + assert self.n_routed_experts % world_size == 0 + global_rank = get_global_rank() + self.num_redundant_experts_per_rank = get_env_start_args().eplb_num_redundant_experts_per_rank + + if self.num_redundant_experts_per_rank > 0: + self.num_primary_experts_per_rank = self.n_routed_experts // world_size + self.num_total_physical_experts = self.n_routed_experts + world_size * self.num_redundant_experts_per_rank + initial_local_expert_ids_by_rank = build_initial_local_expert_ids( + self.n_routed_experts, + world_size, + self.num_redundant_experts_per_rank, + ) + self.local_logics_expert_ids_list = initial_local_expert_ids_by_rank[global_rank] + self.logical_to_physical_map = torch.tensor( + build_logical_to_physical_map( + initial_local_expert_ids_by_rank, + self.n_routed_experts, + current_rank=global_rank, + ), + dtype=torch.int32, + ).cuda() + self.route_counter = torch.zeros(self.n_routed_experts, dtype=torch.int64, device="cuda") + self.recording = True + else: + self.num_total_physical_experts = self.n_routed_experts + num_local_experts = self.n_routed_experts // world_size + first_local_expert_id = global_rank * num_local_experts + self.local_logics_expert_ids_list = list( + range( + first_local_expert_id, + first_local_expert_id + num_local_experts, + ) + ) + def _select_experts( self, input_tensor: torch.Tensor, @@ -40,7 +88,9 @@ def _select_experts( per_expert_scale: Optional[torch.Tensor] = None, ): """Select logical experts without applying the EPLB physical layout.""" - from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts + from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import ( + select_experts, + ) topk_weights, topk_ids = select_experts( hidden_states=input_tensor, @@ -66,15 +116,12 @@ def _prepare_expert_execution( shared_expert_gate: Optional[torch.Tensor] = None, ): assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" - eplb = self.expert_parallel_state.eplb - if eplb is not None: + if self.num_redundant_experts_per_rank > 0: topk_ids = eplb_repair_topk_ids( logical_topk_ids=topk_ids, - logical_to_physical_map=eplb.logical_to_physical_map, - logical_replica_count=eplb.logical_replica_count, - expert_counter=eplb.route_counter, - sample_index=eplb.next_sample_index(), - record_load=eplb.recording, + logical_to_physical_map=self.logical_to_physical_map, + logical_expert_counter=self.route_counter, + update_logical_expert_counter=self.recording, ) return topk_weights, topk_ids @@ -94,7 +141,7 @@ def _fused_experts( w2=w2, topk_weights=topk_weights, topk_idx=topk_ids.to(torch.long), - num_experts=self.expert_parallel_state.num_total_physical_experts, + num_experts=self.num_total_physical_experts, quant_method=self.quant_method, is_prefill=is_prefill, previous_event=None, # for overlap @@ -134,7 +181,7 @@ def low_latency_dispatch( topk_idx=topk_idx, x=hidden_states, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, - num_experts=self.expert_parallel_state.num_total_physical_experts, + num_experts=self.num_total_physical_experts, use_fp8=use_fp8_w8a8, async_finish=False, return_recv_hook=True, @@ -182,7 +229,7 @@ def dispatch( qinput_tensor, topk_idx=topk_idx, topk_weights=topk_weights, - num_experts=self.expert_parallel_state.num_total_physical_experts, + num_experts=self.num_total_physical_experts, num_max_tokens_per_rank=num_max_tokens_per_rank, expert_alignment=128, num_sms=get_ep_num_sms(), @@ -210,7 +257,14 @@ def hook(): compute_load=compute_load, ) - return recv_x, recv_topk_idx, recv_topk_weights, handle.num_recv_tokens_per_expert_list, handle, hook + return ( + recv_x, + recv_topk_idx, + recv_topk_weights, + handle.num_recv_tokens_per_expert_list, + handle, + hook, + ) def masked_group_gemm( self, @@ -293,7 +347,12 @@ def low_latency_combine( handle: Any, ): combined_x, event_overlap, hook = dist_group_manager.ep_low_latency_buffer.low_latency_combine( - gemm_out_b, topk_idx, topk_weights, handle, async_finish=False, return_recv_hook=True + gemm_out_b, + topk_idx, + topk_weights, + handle, + async_finish=False, + return_recv_hook=True, ) return combined_x, hook diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py index f0f491ec65..3bde6cfbc3 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py @@ -4,85 +4,109 @@ @triton.jit -def _replica_index(token_index, logical_id, replica_count): +def _replica_index(token_index, logical_expert_id, num_valid_replicas): token_hash = token_index.to(tl.uint32) * 2654435769 - expert_hash = logical_id.to(tl.uint32) * 2246822519 - return (token_hash + expert_hash) % replica_count.to(tl.uint32) + expert_hash = logical_expert_id.to(tl.uint32) * 2246822519 + return (token_hash + expert_hash) % num_valid_replicas.to(tl.uint32) @triton.jit def _eplb_repair_topk_ids_kernel( logical_topk_ids_ptr, physical_topk_ids_ptr, - topk_id_count, - topk, + num_topk_ids, + top_k, logical_to_physical_map_ptr, - logical_replica_count_ptr, - expert_counter_ptr, - sample_index, - MAP_SLOTS: tl.constexpr, - COUNTER_NUM_EXPERTS: tl.constexpr, - RECORD_LOAD: tl.constexpr, - SINGLE_TOKEN: tl.constexpr, + logical_to_physical_map_row_stride, + logical_expert_counter_ptr, + UPDATE_LOGICAL_EXPERT_COUNTER: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): - offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offsets < topk_id_count - logical_ids = tl.load(logical_topk_ids_ptr + offsets, mask=mask, other=0) + # 阶段 1:将二维 [num_tokens, top_k] 路由结果展平后分块处理。 + # topk_id_offsets 同时用于访问输入、输出,并可恢复它所属的 token 下标。 + topk_id_offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + valid_mask = topk_id_offsets < num_topk_ids + logical_expert_ids = tl.load(logical_topk_ids_ptr + topk_id_offsets, mask=valid_mask, other=0) - if RECORD_LOAD: + # 阶段 2:在 EPLB 采样窗口内,按 logical expert 统计本次路由负载。 + # 统计发生在 physical ID 修复之前,因此冗余副本不会拆散逻辑专家负载。 + if UPDATE_LOGICAL_EXPERT_COUNTER: tl.atomic_add( - expert_counter_ptr + sample_index * COUNTER_NUM_EXPERTS + logical_ids, + logical_expert_counter_ptr + logical_expert_ids, 1, - mask=mask, + mask=valid_mask, sem="relaxed", ) - if SINGLE_TOKEN: - replica_indices = tl.zeros((BLOCK_SIZE,), tl.int32) - else: - replica_counts = tl.load(logical_replica_count_ptr + logical_ids, mask=mask, other=1) - token_indices = offsets // topk - replica_indices = _replica_index(token_indices, logical_ids, replica_counts) + # 阶段 3:定位每个 logical expert 的打包映射行。第 0 列保存有效 + # physical 副本数,第 1 列标记当前 rank 是否有本地副本,第 2 列起 + # 保存可参与路由的 physical expert ID;本地副本存在时固定放在第一个槽。 + map_row_offsets = logical_expert_ids * logical_to_physical_map_row_stride + num_valid_replicas = tl.load(logical_to_physical_map_ptr + map_row_offsets, mask=valid_mask, other=1) + has_local_replica = tl.load(logical_to_physical_map_ptr + map_row_offsets + 1, mask=valid_mask, other=0) - physical_ids = tl.load( - logical_to_physical_map_ptr + logical_ids * MAP_SLOTS + replica_indices, - mask=mask, + # 阶段 4:若当前 rank 持有该专家,强制选择第一个槽位以避免跨 rank + # 通信;否则用 token 下标和 logical expert ID 生成稳定 hash,在有效 + # 副本范围内选择槽位,避免固定 token 位置长期偏向同一个副本。 + token_indices = topk_id_offsets // top_k + hashed_replica_indices = _replica_index(token_indices, logical_expert_ids, num_valid_replicas) + selected_replica_indices = tl.where(has_local_replica != 0, 0, hashed_replica_indices) + + # 阶段 5:读取选中槽位的 physical expert ID 并写入新的输出 tensor。 + # logical_topk_ids 只读,后续 callback 仍可安全观察原始逻辑路由结果。 + physical_expert_ids = tl.load( + logical_to_physical_map_ptr + map_row_offsets + 2 + selected_replica_indices, + mask=valid_mask, other=-1, ) - tl.store(physical_topk_ids_ptr + offsets, physical_ids, mask=mask) + tl.store(physical_topk_ids_ptr + topk_id_offsets, physical_expert_ids, mask=valid_mask) @torch.no_grad() def eplb_repair_topk_ids( logical_topk_ids: torch.Tensor, logical_to_physical_map: torch.Tensor, - logical_replica_count: torch.Tensor, - expert_counter: torch.Tensor, - sample_index: int, - record_load: bool, + logical_expert_counter: torch.Tensor, + update_logical_expert_counter: bool, ) -> torch.Tensor: - """Map logical top-k IDs to the current EPLB physical expert layout.""" + """将 logical top-k ID 转换为当前 EPLB 布局中的 physical expert ID。 + + 参数: + logical_topk_ids: logical expert ID,shape 为 ``[num_tokens, top_k]``。 + logical_to_physical_map: 打包路由表,shape 为 + ``[num_logical_experts, 2 + routing_slots]``。第 0 列保存有效副本 + 数量,第 1 列标记当前 rank 是否有本地副本,其余列保存 physical + expert ID;存在本地副本时,其 physical ID 固定放在第一个槽位。 + 有效副本之后未使用的 padding 槽位为 -1,kernel 不会读取它们。 + logical_expert_counter: 每个 logical expert 的累计路由次数,shape 为 + ``[num_logical_experts]``。 + update_logical_expert_counter: 是否将本次 logical 路由结果累计到 + ``logical_expert_counter``。 + + 返回: + physical expert ID,shape 为 ``[num_tokens, top_k]``。 + """ assert logical_topk_ids.is_contiguous() assert logical_topk_ids.ndim == 2 + assert logical_to_physical_map.ndim == 2 + assert logical_to_physical_map.shape[1] > 2 + assert logical_to_physical_map.stride(1) == 1 + assert logical_expert_counter.ndim == 1 + assert logical_expert_counter.shape[0] == logical_to_physical_map.shape[0] physical_topk_ids = torch.empty_like(logical_topk_ids) if logical_topk_ids.numel() == 0: return physical_topk_ids block_size = 512 _eplb_repair_topk_ids_kernel[(triton.cdiv(logical_topk_ids.numel(), block_size),)]( - logical_topk_ids, - physical_topk_ids, - logical_topk_ids.numel(), - logical_topk_ids.shape[1], - logical_to_physical_map, - logical_replica_count, - expert_counter, - sample_index, - MAP_SLOTS=logical_to_physical_map.shape[1], - COUNTER_NUM_EXPERTS=expert_counter.shape[1], - RECORD_LOAD=record_load, - SINGLE_TOKEN=logical_topk_ids.shape[0] == 1, + logical_topk_ids_ptr=logical_topk_ids, + physical_topk_ids_ptr=physical_topk_ids, + num_topk_ids=logical_topk_ids.numel(), + top_k=logical_topk_ids.shape[1], + logical_to_physical_map_ptr=logical_to_physical_map, + logical_to_physical_map_row_stride=logical_to_physical_map.stride(0), + logical_expert_counter_ptr=logical_expert_counter, + UPDATE_LOGICAL_EXPERT_COUNTER=update_logical_expert_counter, BLOCK_SIZE=block_size, num_warps=4, num_stages=1, diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 081708083c..9da7cec8fb 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -50,9 +50,6 @@ ) from .multi_level_kv_cache import MultiLevelKvCacheModule from lightllm.utils.profiler import ProcessProfiler, ProfilerCmd -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - disable_eplb_model_init, -) class ModeBackend: @@ -350,8 +347,7 @@ def init_mtp_draft_model(self, main_kvargs: dict): model_cfg=draft_model_cfg, spec_mode=spec_mode, ) - with disable_eplb_model_init(): - draft_model = draft_model_class(draft_model_kvargs) + draft_model = draft_model_class(draft_model_kvargs) self.draft_models.append(draft_model) self.logger.info(f"loaded speculative draft model class {self.draft_models[i].__class__}") diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 456565e81d..e411a6abdb 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -6,8 +6,11 @@ import torch.distributed as dist from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import ( + FusedMoeWeight, +) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_initial_local_expert_ids, build_logical_to_physical_maps_for_layers, plan_redundant_experts, select_improving_placements, @@ -17,7 +20,11 @@ align_target_placement, build_transfer_plan, ) -from lightllm.utils.dist_utils import get_global_rank, get_global_world_size, get_node_world_size +from lightllm.utils.dist_utils import ( + get_global_rank, + get_global_world_size, + get_node_world_size, +) from lightllm.utils.envs_utils import ( get_eplb_placement_stickiness, get_eplb_rebalance_gain_threshold, @@ -41,19 +48,29 @@ def __init__(self, model: TpPartBaseModel): self.global_rank = get_global_rank() self.world_size = get_global_world_size() self.node_world_size = get_node_world_size() - self._eplb_states = [weight.expert_parallel_state.eplb for weight in self.weights] + self._eplb_impls = [weight.fuse_moe_impl for weight in self.weights] self.step_interval = get_prefill_eplb_step_interval() self.rebalance_gain_threshold = get_eplb_rebalance_gain_threshold() self.placement_stickiness = get_eplb_placement_stickiness() self.sampling_interval = self.step_interval self.prefill_steps = 0 - routed = {weight.expert_parallel_state.num_logical_experts for weight in self.weights} - redundant = {state.num_redundant_experts_per_rank for state in self._eplb_states} + routed = {weight.fuse_moe_impl.n_routed_experts for weight in self.weights} + redundant = {impl.num_redundant_experts_per_rank for impl in self._eplb_impls} assert len(routed) == len(redundant) == 1 self.num_logical_experts = routed.pop() self.num_redundant_experts_per_rank = redundant.pop() - self.current_placement = torch.stack( - [state.initial_redundant_expert_ids_by_rank for state in self._eplb_states] + num_primary_experts_per_rank = self.num_logical_experts // self.world_size + initial_local_expert_ids_by_rank = build_initial_local_expert_ids( + self.num_logical_experts, + self.world_size, + self.num_redundant_experts_per_rank, + ) + initial_redundant_expert_ids_by_rank = [ + expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids_by_rank + ] + self.current_placement = torch.tensor( + [initial_redundant_expert_ids_by_rank for _ in self.weights], + dtype=torch.int64, ) self.in_flight = False self.target_placement = None @@ -71,7 +88,7 @@ def __init__(self, model: TpPartBaseModel): self._continuous_collection_start_step: Optional[int] = None self._continuous_collection_end_step: Optional[int] = self.step_interval self._steady_collection_end_step: Optional[int] = None - self._reset_recorded_samples() + self._reset_route_counters() self._set_recording(True) # Keep background evaluation collectives separate from the main-thread # control/poll collectives: their ordering is intentionally independent. @@ -127,15 +144,13 @@ def step(self): self._arm_steady_collection(self.prefill_steps + self._steady_sample_window_steps()) def _set_recording(self, enabled: bool): - for state in self._eplb_states: - state.recording = enabled + for impl in self._eplb_impls: + impl.recording = enabled - def _reset_recorded_samples(self): - counters = [state.route_counter for state in self._eplb_states] + def _reset_route_counters(self): + counters = [impl.route_counter for impl in self._eplb_impls] if counters: torch._foreach_zero_(counters) - for state in self._eplb_states: - state.recorded_sample_count = 0 def _control_count(self, value: int) -> torch.Tensor: """Return the main-thread-only reusable control collective scalar.""" @@ -150,14 +165,14 @@ def _steady_sample_window_steps(self) -> int: def _arm_steady_collection(self, collection_end_step: int): """Start the fixed sparse window without moving its evaluation boundary.""" - self._reset_recorded_samples() + self._reset_route_counters() self._steady_collection_end_step = collection_end_step self._set_recording(True) def _begin_continuous_collection(self): minimum_end = self.prefill_steps + self.step_interval collection_end = -(-minimum_end // self.sampling_interval) * self.sampling_interval - self._reset_recorded_samples() + self._reset_route_counters() self._steady_collection_end_step = None self._continuous_collection_start_step = collection_end - self.step_interval self._continuous_collection_end_step = collection_end @@ -168,7 +183,7 @@ def _prepare_next_sampling_window(self): self._clear_continuous_collection() self._steady_collection_end_step = None if self.sampling_interval == 1: - self._reset_recorded_samples() + self._reset_route_counters() self._set_recording(True) elif self.sampling_interval <= EPLB_STEADY_SAMPLE_STEPS: # There is no later pre-boundary manager step at which to arm a @@ -176,45 +191,18 @@ def _prepare_next_sampling_window(self): # fixed boundary. self._arm_steady_collection(self.prefill_steps + self.sampling_interval) else: - self._reset_recorded_samples() + self._reset_route_counters() self._set_recording(False) - @staticmethod - def _recent_ring_samples(counter: torch.Tensor, recorded_sample_count: int) -> torch.Tensor: - """Return the newest ring rows in chronological order.""" - capacity = counter.shape[0] - available = min(recorded_sample_count, capacity) - if available == 0: - return counter[:0] - start = (recorded_sample_count - available) % capacity - indices = (torch.arange(available, dtype=torch.int64, device=counter.device) + start) % capacity - return counter.index_select(0, indices) - def _collect_local_samples(self) -> torch.Tensor: - counters = [state.route_counter for state in self._eplb_states] - capacities = [counter.shape[0] for counter in counters] - if len(set(capacities)) != 1 or any(counter.ndim != 2 for counter in counters): - raise RuntimeError("EPLB sample capacities differ between layers") - counts = [state.recorded_sample_count for state in self._eplb_states] - if len(set(counts)) != 1: - raise RuntimeError("EPLB recorded sample counts differ between layers") - sample_count = counts[0] - # Validate the metadata before copying the newest rows to the CPU. - metadata = torch.tensor([sample_count, -sample_count, capacities[0], -capacities[0]], dtype=torch.int64) - dist.all_reduce(metadata, op=dist.ReduceOp.MIN, group=self.evaluation_group) - if metadata[0] != -metadata[1] or metadata[2] != -metadata[3]: - raise RuntimeError("EPLB recorded sample count or capacity differs between ranks") - # Stack the fixed-size ring buffers in one GPU launch. Slicing each - # layer before stacking turns a single launch into one index_select per - # MoE layer and is measurably slower in the normal sparse path. - counter_samples = torch.stack(counters, dim=1) - return self._recent_ring_samples(counter_samples, sample_count).cpu() + counters = [impl.route_counter for impl in self._eplb_impls] + if any(counter.ndim != 1 or counter.shape[0] != self.num_logical_experts for counter in counters): + raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") + return torch.stack(counters).unsqueeze(0).cpu() def _commit_layer_metadata(self, layer_index: int): - eplb_state = self._eplb_states[layer_index] - logical_to_physical, replica_count = self.target_metadata[layer_index] - eplb_state.logical_to_physical_map.copy_(logical_to_physical, non_blocking=True) - eplb_state.logical_replica_count.copy_(replica_count, non_blocking=True) + impl = self._eplb_impls[layer_index] + impl.logical_to_physical_map.copy_(self.target_metadata[layer_index], non_blocking=True) def _finish_rebalance(self): self.current_placement = self.target_placement @@ -253,7 +241,11 @@ def _poll_in_flight(self): raise RuntimeError( f"EPLB pending layer {layer_index} does not match expected {self.in_flight_layers[0]}" ) - self.transfer.commit(layer_index, buffer_index, lambda: self._commit_layer_metadata(layer_index)) + self.transfer.commit( + layer_index, + buffer_index, + lambda: self._commit_layer_metadata(layer_index), + ) self.in_flight_layers.pop(0) if not self.in_flight_layers: self.transfer.finish() @@ -279,7 +271,6 @@ def _plan_and_broadcast(self, global_load: torch.Tensor): self.world_size, self.num_redundant_experts_per_rank, expert_alignment=EPLB_EXPERT_ALIGNMENT, - node_world_size=self.node_world_size, current_placement=self.current_placement, stickiness=self.placement_stickiness, ) @@ -288,7 +279,6 @@ def _plan_and_broadcast(self, global_load: torch.Tensor): self.current_placement, candidate, expert_alignment=EPLB_EXPERT_ALIGNMENT, - node_world_size=self.node_world_size, rebalance_gain_threshold=self.rebalance_gain_threshold, ) if bool(torch.any(improved)): @@ -300,10 +290,11 @@ def _plan_and_broadcast(self, global_load: torch.Tensor): placement = placement.clone() for layer_index in torch.nonzero(improved, as_tuple=False).flatten().tolist(): placement[layer_index] = align_target_placement( - self.current_placement[layer_index], placement[layer_index] + self.current_placement[layer_index], + placement[layer_index], ) result = { - "kind": "planned" if bool(torch.any(improved)) else "no_improvement", + "kind": ("planned" if bool(torch.any(improved)) else "no_improvement"), "placement": placement, "improved": improved, "before": _imbalance_summary(before_load), @@ -326,42 +317,55 @@ def _plan_and_broadcast(self, global_load: torch.Tensor): def _evaluate_after_event(self, event: torch.cuda.Event): """Run the CPU/Gloo planning phase after the frozen CUDA counters are ready.""" try: - torch.cuda.set_device(self._eplb_states[0].route_counter.device) + torch.cuda.set_device(self._eplb_impls[0].route_counter.device) event.synchronize() local_load = self._collect_local_samples() - recorded_sample_count = int(local_load.shape[0]) sample_window_steps = ( self.step_interval if self._continuous_collection_end_step is not None else self._steady_sample_window_steps() ) - num_nodes = self.world_size // self.node_world_size - # Preserve source nodes until physical-replica loads are combined; - # DeepEP applies expert alignment after traffic from all sources - # reaches each destination expert. - global_load = torch.zeros((*local_load.shape[:2], num_nodes, local_load.shape[2]), dtype=local_load.dtype) - global_load[:, :, self.global_rank // self.node_world_size] = local_load + # 保留每个当前 rank 的负载,planner 才能准确模拟“本卡优先, + # 否则在所有远端副本间分配”的运行时路由规则。 + global_load = torch.zeros( + (*local_load.shape[:2], self.world_size, local_load.shape[2]), + dtype=local_load.dtype, + ) + global_load[:, :, self.global_rank] = local_load dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) result = self._plan_and_broadcast(global_load) - result["recorded_sample_count"] = recorded_sample_count result["sample_window_steps"] = sample_window_steps if result["kind"] == "planned": metadata = [None] * len(self.weights) layer_plans = [] improved_layer_indices = torch.nonzero(result["improved"], as_tuple=False).flatten() if improved_layer_indices.numel(): - maps_for_improved_layers, counts_for_improved_layers = build_logical_to_physical_maps_for_layers( - result["placement"][improved_layer_indices], - self.num_logical_experts, - source_rank=self.global_rank, - node_world_size=self.node_world_size, + redundant_placements = result["placement"][improved_layer_indices].tolist() + num_primary_experts_per_rank = self.num_logical_experts // self.world_size + rank_to_logic_expert_ids_by_layer = [ + [ + list( + range( + rank * num_primary_experts_per_rank, + (rank + 1) * num_primary_experts_per_rank, + ) + ) + + rank_redundant_expert_ids + for rank, rank_redundant_expert_ids in enumerate(layer_placement) + ] + for layer_placement in redundant_placements + ] + maps_for_improved_layers = torch.tensor( + build_logical_to_physical_maps_for_layers( + rank_to_logic_expert_ids_by_layer, + self.num_logical_experts, + current_rank=self.global_rank, + ), + dtype=torch.int32, ) for improved_layer_offset, layer_index in enumerate(improved_layer_indices.tolist()): placement = result["placement"][layer_index] - metadata[layer_index] = ( - maps_for_improved_layers[improved_layer_offset], - counts_for_improved_layers[improved_layer_offset], - ) + metadata[layer_index] = maps_for_improved_layers[improved_layer_offset] layer_plans.append( ( layer_index, @@ -424,25 +428,23 @@ def _poll_evaluation(self): if from_continuous_window: logger.info( "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " - "next_sampling_interval=%s recorded_sample_count=%s sample_window_steps=%s", + "next_sampling_interval=%s sample_window_steps=%s", self.prefill_steps, result["minimum_layer_samples"], result["minimum"], self.sampling_interval, - result.get("recorded_sample_count"), result.get("sample_window_steps"), ) else: logger.info( "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " "scheduled_fresh_window_start=%s scheduled_fresh_window_end=%s " - "recorded_sample_count=%s sample_window_steps=%s", + "sample_window_steps=%s", self.prefill_steps, result["minimum_layer_samples"], result["minimum"], self._continuous_collection_start_step, self._continuous_collection_end_step, - result.get("recorded_sample_count"), result.get("sample_window_steps"), ) return False @@ -453,13 +455,12 @@ def _poll_evaluation(self): "eplb skip rearrangement: no model improvement model_imbalance_ratio=%.4f " "candidate_model_imbalance_ratio=%.4f candidate_rebalance_gain=%.4f " "candidate_changed_layer_count=%s actual_changed_layer_count=0 next_sampling_interval=%s " - "recorded_sample_count=%s sample_window_steps=%s", + "sample_window_steps=%s", result["model_imbalance_ratio"], result["candidate_model_imbalance_ratio"], result["candidate_rebalance_gain"], result["candidate_changed_layer_count"], self.sampling_interval, - result.get("recorded_sample_count"), result.get("sample_window_steps"), ) self._prepare_next_sampling_window() @@ -486,7 +487,7 @@ def _start_rebalance(self, result): layer_plans = result["layer_plans"] self.sampling_interval = self.step_interval self._clear_continuous_collection() - self._reset_recorded_samples() + self._reset_route_counters() self.target_placement = placement self.target_metadata = result["metadata"] self.in_flight_layers = [layer_index for layer_index, _ in layer_plans] @@ -505,7 +506,7 @@ def _start_rebalance(self, result): "model_imbalance_ratio=%.4f candidate_model_imbalance_ratio=%.4f " "candidate_rebalance_gain=%.4f candidate_changed_layer_count=%s " "actual_changed_layer_count=%s actual_changed_slot_count=%s cross_node_transfer_count=%s " - "recorded_sample_count=%s sample_window_steps=%s", + "sample_window_steps=%s", self.prefill_steps, result["before"]["max"], result["after"]["max"], @@ -518,7 +519,6 @@ def _start_rebalance(self, result): len(layer_plans), actual_changed_slot_count, cross_node_transfer_count, - result.get("recorded_sample_count"), result.get("sample_window_steps"), ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index e5737fcb7a..eb0ac23fa3 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -4,7 +4,7 @@ import re import socket import threading -from collections import defaultdict, deque +from collections import Counter, defaultdict, deque from dataclasses import dataclass from typing import Dict, Iterable, List, Optional, Sequence, Tuple @@ -38,11 +38,19 @@ def align_target_placement(current: torch.Tensor, target: torch.Tensor) -> torch target_rows = target.tolist() aligned_target_rows = [] for current_row, target_row in zip(current_rows, target_rows): - current_slots = {expert: slot for slot, expert in enumerate(current_row)} - target_experts = set(target_row) + remaining_target = Counter(target_row) aligned_row = list(current_row) - freed_slots = [slot for slot, expert in enumerate(current_row) if expert not in target_experts] - new_experts = [expert for expert in target_row if expert not in current_slots] + freed_slots = [] + for slot, expert in enumerate(current_row): + if remaining_target[expert] > 0: + remaining_target[expert] -= 1 + else: + freed_slots.append(slot) + new_experts = [] + for expert in target_row: + if remaining_target[expert] > 0: + new_experts.append(expert) + remaining_target[expert] -= 1 assert len(freed_slots) == len(new_experts) for slot, expert in zip(freed_slots, new_experts): aligned_row[slot] = expert @@ -61,9 +69,8 @@ def build_transfer_plan( num_experts_per_rank = num_logical_experts // world_size current_rows = current.tolist() aligned_target_rows = align_target_placement(current, target).tolist() - # A logical expert has one primary row and at most one redundant row per - # rank, so this source list is already unique. Build it once instead of - # allocating/sorting a set for every destination slot. + # 初始占位布局允许同一个 logical expert 在一个 rank 上出现多次,因此 + # 保留所有物理来源候选,让首次迁移也可以从任意已加载的副本复制。 candidates_by_expert = [ [ ( @@ -100,15 +107,15 @@ class _EPLBTransferBase: """Shared live/staging buffers and publish/commit lifecycle.""" def __init__(self, weights, transfer_group, global_rank, world_size): - self._eplb_states = [weight.expert_parallel_state.eplb for weight in weights] + self._eplb_impls = [weight.fuse_moe_impl for weight in weights] self.transfer_group = transfer_group self.global_rank = global_rank self.world_size = world_size - self.num_experts_per_rank = weights[0].expert_parallel_state.num_primary_experts_per_rank + self.num_experts_per_rank = weights[0].fuse_moe_impl.num_primary_experts_per_rank self.device = weights[0].w13.weight.device self.live = [extract_eplb_expert_tensors(weight) for weight in weights] self._validate_live_layout() - num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank + num_redundant_slots_per_rank = self._eplb_impls[0].num_redundant_experts_per_rank self.staging = [ [ ( @@ -137,12 +144,12 @@ def __init__(self, weights, transfer_group, global_rank, world_size): def _validate_live_layout(self) -> None: reference = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in self.live[0]] - num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank - for layer_index, (state, tensors) in enumerate(zip(self._eplb_states, self.live)): + num_redundant_slots_per_rank = self._eplb_impls[0].num_redundant_experts_per_rank + for layer_index, (impl, tensors) in enumerate(zip(self._eplb_impls, self.live)): layout = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in tensors] assert layout == reference, f"EPLB layer {layer_index} has incompatible expert tensor layout" assert ( - state.num_redundant_experts_per_rank == num_redundant_slots_per_rank + impl.num_redundant_experts_per_rank == num_redundant_slots_per_rank ), "EPLB redundant slot count must match" def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index afa48c094c..7ca506f767 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -10,7 +10,7 @@ import torch from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_initial_redundant_expert_ids, + build_initial_local_expert_ids, build_logical_to_physical_map, build_logical_to_physical_maps_for_layers, _estimate_rank_load, @@ -29,10 +29,6 @@ from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( deepgemm_impl as deepgemm_module, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - EPLBState, - ExpertParallelState, -) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( create_fuse_moe_impl, FuseMoeMarlin, @@ -44,10 +40,6 @@ from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe import ( fused_moe_weight as fused_weight_module, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - disable_eplb_model_init, - is_eplb_model_init_disabled, -) from lightllm.common.eplb_utils import extract_eplb_expert_tensors from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( TransferStep, @@ -58,7 +50,7 @@ ) -def _test_parallel_state( +def _test_moe_impl( *, eplb=False, num_logical_experts=128, @@ -66,85 +58,89 @@ def _test_parallel_state( num_redundant_experts_per_rank=1, route_counter=None, recording=False, - recorded_sample_count=0, ): - eplb_state = None + logical_to_physical_map = None if eplb: if route_counter is None: - route_counter = torch.zeros((2, num_logical_experts), dtype=torch.int64) - initial_layout_world_size = max(world_size, 2) - eplb_state = EPLBState( - num_redundant_experts_per_rank=num_redundant_experts_per_rank, - initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids( - num_logical_experts, - initial_layout_world_size, - num_redundant_experts_per_rank, - ), - logical_to_physical_map=torch.zeros((num_logical_experts, 2), dtype=torch.int32), - logical_replica_count=torch.ones(num_logical_experts, dtype=torch.int32), - route_counter=route_counter, - recording=recording, - recorded_sample_count=recorded_sample_count, - ) - return ExpertParallelState( - num_logical_experts=num_logical_experts, - world_size=world_size, - eplb=eplb_state, + route_counter = torch.zeros((num_logical_experts,), dtype=torch.int64) + logical_to_physical_map = torch.zeros((num_logical_experts, world_size + 2), dtype=torch.int32) + logical_to_physical_map[:, 0] = 1 + else: + num_redundant_experts_per_rank = 0 + return SimpleNamespace( + n_routed_experts=num_logical_experts, + num_primary_experts_per_rank=num_logical_experts // world_size, + num_total_physical_experts=(num_logical_experts + world_size * num_redundant_experts_per_rank), + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + logical_to_physical_map=logical_to_physical_map, + route_counter=route_counter, + recording=recording, + ) + + +def _set_deepgemm_runtime(impl, runtime): + for name in ( + "num_primary_experts_per_rank", + "num_total_physical_experts", + "num_redundant_experts_per_rank", + "logical_to_physical_map", + "route_counter", + "recording", + ): + setattr(impl, name, getattr(runtime, name)) + + +def _initial_extra_expert_placement(num_logical_experts, world_size, num_redundant_experts_per_rank): + num_primary_experts_per_rank = num_logical_experts // world_size + initial_local_expert_ids_by_rank = build_initial_local_expert_ids( + num_logical_experts, world_size, num_redundant_experts_per_rank + ) + return torch.tensor( + [expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids_by_rank], + dtype=torch.int64, ) -def _validated_expert_parallel_state( - *, - eplb=True, - n_routed_experts=4, - world_size=2, - num_redundant_experts_per_rank=1, - device="cpu", -): - runtime = None - if eplb: - runtime = EPLBState( - num_redundant_experts_per_rank=num_redundant_experts_per_rank, - initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids( - n_routed_experts, - world_size, - num_redundant_experts_per_rank, - ), - logical_to_physical_map=torch.empty((n_routed_experts, 2), dtype=torch.int32, device=device), - logical_replica_count=torch.ones(n_routed_experts, dtype=torch.int32, device=device), - route_counter=torch.zeros((2, n_routed_experts), dtype=torch.int64, device=device), +def _rank_to_logic_expert_ids(redundant_placement, num_logical_experts): + num_ranks = len(redundant_placement) + num_primary_experts_per_rank = num_logical_experts // num_ranks + return [ + list( + range( + rank * num_primary_experts_per_rank, + (rank + 1) * num_primary_experts_per_rank, + ) ) - return ExpertParallelState( - num_logical_experts=n_routed_experts, - world_size=world_size, - eplb=runtime, - ) - - -def _set_expert_parallel_state(impl, state): - impl.expert_parallel_state = state + + list(rank_redundant_expert_ids) + for rank, rank_redundant_expert_ids in enumerate(redundant_placement) + ] -def _manual_runtime_rank_load(source_load, placement, node_world_size, alignment): +def _manual_runtime_rank_load(source_load, placement, alignment): """Reference the committed runtime logical-to-physical maps on CPU.""" - samples, layers, nodes, num_logical_experts = source_load.shape + samples, layers, source_ranks, num_logical_experts = source_load.shape ranks, redundant = placement.shape[1:] + assert source_ranks == ranks num_experts_per_rank = num_logical_experts // ranks num_physical_experts_per_rank = num_experts_per_rank + redundant raw = torch.zeros((samples, layers, ranks, num_logical_experts), dtype=torch.float64) for layer in range(layers): - for source_node in range(nodes): - logical_to_physical, replica_count = build_logical_to_physical_map( - placement[layer], + for current_rank in range(ranks): + logical_to_physical = build_logical_to_physical_map( + _rank_to_logic_expert_ids(placement[layer].tolist(), num_logical_experts), num_logical_experts, - source_rank=source_node * node_world_size, - node_world_size=node_world_size, + current_rank=current_rank, ) - for expert in range(num_logical_experts): - count = int(replica_count[expert].item()) - for physical_id in logical_to_physical[expert, :count].tolist(): + for expert, packed_row in enumerate(logical_to_physical): + num_replicas = packed_row[0] + physical_expert_ids = packed_row[2 : num_replicas + 2] + if packed_row[1]: + physical_expert_ids = physical_expert_ids[:1] + for physical_id in physical_expert_ids: rank = physical_id // num_physical_experts_per_rank - raw[:, layer, rank, expert] += source_load[:, layer, source_node, expert] / count + raw[:, layer, rank, expert] += source_load[:, layer, current_rank, expert] / len( + physical_expert_ids + ) return (torch.ceil(raw / alignment) * alignment).sum(dim=3) @@ -204,27 +200,40 @@ def _fused_experts( assert seen["fused"]["topk_ids"] == "physical_ids" -def test_parallel_state_derives_expert_layout(): - state = _validated_expert_parallel_state(eplb=True) - assert state.num_primary_experts_per_rank == 2 - assert state.num_total_physical_experts == 6 +def test_deepgemm_runtime_derives_expert_layout(): + runtime = _test_moe_impl( + eplb=True, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + assert runtime.num_primary_experts_per_rank == 2 + assert runtime.num_total_physical_experts == 6 -def test_factory_selects_all_paths_and_requires_ep_state(monkeypatch): +def test_factory_selects_all_paths_without_ep_constructor_state(monkeypatch): plain_quant = SimpleNamespace(method_name="none") marlin_quant = SimpleNamespace(method_name="awq_marlin") monkeypatch.setattr(FuseMoeMarlin, "create_workspace", lambda self: None) - state = _validated_expert_parallel_state(eplb=False) + monkeypatch.setattr( + deepgemm_module, + "get_env_start_args", + lambda: SimpleNamespace(eplb_num_redundant_experts_per_rank=0), + ) + monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) ep_impl = create_fuse_moe_impl( n_routed_experts=4, num_fused_shared_experts=0, routed_scaling_factor=1.0, quant_method=plain_quant, - expert_parallel_state=state, + enable_ep_moe=True, ) assert isinstance(ep_impl, deepgemm_module.FuseMoeDeepGEMM) - assert ep_impl.expert_parallel_state is state - assert state.eplb is None + assert ep_impl.num_total_physical_experts == 4 + assert not hasattr(ep_impl, "num_primary_experts_per_rank") + assert not hasattr(ep_impl, "route_counter") + assert not hasattr(ep_impl, "expert_parallel_state") assert isinstance( create_fuse_moe_impl( n_routed_experts=4, @@ -279,25 +288,23 @@ def test_eplb_redundant_experts_default_to_disabled(): @pytest.mark.parametrize( ("num_logical_experts", "num_ranks", "num_redundant_experts_per_rank", "expected"), [ - (8, 4, 2, [[2, 3], [4, 5], [6, 7], [0, 1]]), - (6, 3, 4, [[2, 3, 4, 5], [4, 5, 0, 1], [0, 1, 2, 3]]), + (8, 4, 2, [[0, 1, 1, 1], [2, 3, 3, 3], [4, 5, 5, 5], [6, 7, 7, 7]]), + (6, 3, 4, [[0, 1, 1, 1, 1, 1], [2, 3, 3, 3, 3, 3], [4, 5, 5, 5, 5, 5]]), ], ) -def test_build_initial_redundant_expert_ids( +def test_build_initial_local_expert_ids( num_logical_experts, num_ranks, num_redundant_experts_per_rank, expected, ): - actual = build_initial_redundant_expert_ids( + actual = build_initial_local_expert_ids( num_logical_experts, num_ranks, num_redundant_experts_per_rank, ) - assert actual.dtype == torch.int64 - assert actual.shape == (num_ranks, num_redundant_experts_per_rank) - assert torch.equal(actual, torch.tensor(expected, dtype=torch.int64)) + assert actual == expected def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): @@ -532,7 +539,11 @@ def test_select_improving_placements_accepts_five_percent_critical_reduction(): candidate = torch.tensor([[[2], [1]]]) selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - samples, current, candidate, rebalance_gain_threshold=0.05, expert_alignment=128 + samples, + current, + candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, ) assert improved.item() @@ -569,69 +580,98 @@ def test_select_improving_placements_accepts_only_when_model_gain_reaches_thresh assert torch.equal(selected, candidate) -def test_logical_to_physical_map_has_at_most_one_slot_per_rank(): - redundant_expert_ids = torch.tensor([[2, 3], [0, 1]]) - logical_to_physical, replica_count = build_logical_to_physical_map(redundant_expert_ids, num_logical_experts=4) +def test_logical_to_physical_map_selects_one_physical_expert(): + rank_to_logic_expert_ids = [[0, 1, 2, 3], [2, 3, 0, 1]] + logical_to_physical = build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=4, current_rank=0) - assert logical_to_physical.shape == (4, 2) - assert torch.equal(replica_count, torch.full((4,), 2, dtype=torch.int64)) + assert isinstance(logical_to_physical, list) + assert len(logical_to_physical) == 4 + assert all(len(row) == 7 for row in logical_to_physical) + assert [row[0] for row in logical_to_physical] == [2, 2, 2, 2] + assert [row[1] for row in logical_to_physical] == [1, 1, 1, 1] + assert [row[2] for row in logical_to_physical] == [0, 1, 2, 3] + assert all(physical_id >= 0 for row in logical_to_physical for physical_id in row[2 : 2 + row[0]]) + assert all(physical_id == -1 for row in logical_to_physical for physical_id in row[2 + row[0] :]) -def test_logical_to_physical_map_requires_expert_count_divisible_by_rank_count_without_source_rank(): - redundant_expert_ids = torch.tensor([[0], [1]]) +def test_logical_to_physical_map_requires_expert_count_divisible_by_rank_count(): + rank_to_logic_expert_ids = [[0, 1, 0], [2, 3, 1]] with pytest.raises(AssertionError): - build_logical_to_physical_map(redundant_expert_ids, num_logical_experts=5) + build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=5, current_rank=0) -def test_logical_to_physical_map_prefers_source_node_replicas(): - # Four ranks, two ranks per node, two primary experts/rank and one - # redundant slot/rank. Expert 0 is primary on rank 0 and replicated on - # rank 2 (the other node). - redundant = torch.tensor([[4], [5], [0], [1]], dtype=torch.int64) - rank0_map, rank0_count = build_logical_to_physical_map( - redundant, num_logical_experts=8, source_rank=0, node_world_size=2 - ) - rank1_map, rank1_count = build_logical_to_physical_map( - redundant, num_logical_experts=8, source_rank=1, node_world_size=2 +def test_logical_to_physical_map_supports_all_redundant_slots_for_one_expert(): + logical_to_physical = build_logical_to_physical_map( + [[0, 1, 0, 0], [2, 3, 0, 0]], + num_logical_experts=4, + current_rank=0, ) - fallback_redundant = torch.tensor([[4], [5], [0], [1], [2], [3]], dtype=torch.int64) - rank4_map, rank4_count = build_logical_to_physical_map(fallback_redundant, 12, source_rank=4, node_world_size=2) - assert rank0_count[0].item() == rank1_count[0].item() == 1 - assert torch.equal(rank0_map[0, :1], torch.tensor([0])) - assert torch.equal(rank1_map[0, :1], torch.tensor([0])) - assert rank4_count[0].item() == 2 - assert set(rank4_map[0, :2].tolist()) == {0, 8} + # 1 个主副本加上 2 个 rank 的全部 4 个冗余槽。 + assert logical_to_physical[0][0] == 5 + assert len(logical_to_physical[0][2:]) == 5 + assert len(set(logical_to_physical[0][2:])) == 5 + + +def test_logical_to_physical_map_prefers_current_rank_replica(): + redundant = [[4], [5], [0], [1]] + rank_to_logic_expert_ids = _rank_to_logic_expert_ids(redundant, 8) + rank0_map = build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=8, current_rank=0) + rank1_map = build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=8, current_rank=1) + fallback_redundant = [[4], [5], [0], [1], [2], [3]] + rank4_map = build_logical_to_physical_map(_rank_to_logic_expert_ids(fallback_redundant, 12), 12, current_rank=4) + + assert rank0_map[0][0] == rank1_map[0][0] == 2 + assert rank0_map[0][1] == 1 + assert rank1_map[0][1] == 0 + assert rank0_map[0][2] == 0 + assert set(rank0_map[0][2:4]) == {0, 8} + assert set(rank1_map[0][2:4]) == {0, 8} + assert rank4_map[0][0] == 2 + assert rank4_map[0][1] == 0 + assert set(rank4_map[0][2:4]) == {0, 8} + + +def test_nonlocal_rank_routes_across_all_replicas(): + redundant = [[4], [5], [0], [1]] + rank_to_logic_expert_ids = _rank_to_logic_expert_ids(redundant, 8) + maps = [build_logical_to_physical_map(rank_to_logic_expert_ids, 8, current_rank=rank) for rank in range(4)] + assert maps[0][0][1] == 1 + assert maps[1][0][1] == 0 + assert maps[2][0][1] == 1 + assert maps[3][0][1] == 0 + assert all(set(logical_map[0][2:4]) == {0, 8} for logical_map in maps) + +def test_current_rank_moves_local_replica_to_front_without_changing_copies(): + # Expert 0 is primary on rank 0 and redundant on rank 1。两个 rank 都优先 + # 自己的本地副本,因此路由槽的起点不同,但候选集合和数量保持一致。 + redundant = [[1], [0], [3], [2]] + rank_to_logic_expert_ids = _rank_to_logic_expert_ids(redundant, 4) + rank0_map = build_logical_to_physical_map(rank_to_logic_expert_ids, 4, current_rank=0) + rank1_map = build_logical_to_physical_map(rank_to_logic_expert_ids, 4, current_rank=1) -def test_source_node_local_maps_fall_back_to_global_replicas(): - redundant = torch.tensor([[4], [5], [0], [1]], dtype=torch.int64) - maps = [build_logical_to_physical_map(redundant, 8, source_rank=rank, node_world_size=2) for rank in range(4)] - assert maps[0][1][0].item() == maps[1][1][0].item() == 1 - assert maps[2][1][0].item() == maps[3][1][0].item() == 1 - assert maps[0][0][0, 0].item() == maps[1][0][0, 0].item() == 0 - assert maps[2][0][0, 0].item() == maps[3][0][0, 0].item() == 8 + assert rank0_map[0] == [2, 1, 0, 3, -1, -1, -1] + assert rank1_map[0] == [2, 1, 3, 0, -1, -1, -1] -def test_source_rank_rotates_selected_replica_order_without_changing_copies(): - # Expert 0 is primary on rank 0 and redundant on rank 1, so both ranks - # on node 0 have the same two local copies. Their source-rank phases - # must differ while their selected set/count remain identical. - redundant = torch.tensor([[1], [0], [3], [2]], dtype=torch.int64) - rank0_map, rank0_count = build_logical_to_physical_map(redundant, 4, source_rank=0, node_world_size=2) - rank1_map, rank1_count = build_logical_to_physical_map(redundant, 4, source_rank=1, node_world_size=2) +def test_current_rank_stably_moves_all_local_physical_ids_to_front(): + # Expert 0 在 rank 0/1/2 上依次对应 physical IDs [0, 1, 3, 5]。 + # 对 rank 1 构建路由表时,只把本地 ID 3 移到最前面;其余远端 ID + # 仍保持原来的 [0, 1, 5] 顺序。 + rank_to_logic_expert_ids = [[0, 0], [1, 0], [2, 0]] - assert rank0_count[0].item() == rank1_count[0].item() == 2 - assert set(rank0_map[0, :2].tolist()) == set(rank1_map[0, :2].tolist()) == {0, 3} - assert torch.equal(rank1_map[0, :2], torch.tensor([3, 0], dtype=torch.int32)) + rank1_map = build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=3, current_rank=1) + + assert rank1_map[0] == [4, 1, 3, 0, 1, 5] @pytest.mark.parametrize( - "source_rank,node_world_size", - [(None, None), (0, 2), (1, 2), (2, 2), (3, 2)], + "current_rank", + [0, 1, 2, 3], ) -def test_logical_to_physical_maps_for_layers_match_single_layer_api(source_rank, node_world_size): +def test_logical_to_physical_maps_for_layers_match_single_layer_api(current_rank): placements_by_layer = torch.tensor( [ [[4, 5], [0, 1], [0, 1], [2, 3]], @@ -641,82 +681,76 @@ def test_logical_to_physical_maps_for_layers_match_single_layer_api(source_rank, dtype=torch.int64, ) - maps_by_layer, counts_by_layer = build_logical_to_physical_maps_for_layers( - placements_by_layer, + rank_to_logic_expert_ids_by_layer = [ + _rank_to_logic_expert_ids(placement.tolist(), 8) for placement in placements_by_layer + ] + maps_by_layer = build_logical_to_physical_maps_for_layers( + rank_to_logic_expert_ids_by_layer, num_logical_experts=8, - source_rank=source_rank, - node_world_size=node_world_size, + current_rank=current_rank, ) expected_by_layer = [ build_logical_to_physical_map( - placement, + rank_to_logic_expert_ids, num_logical_experts=8, - source_rank=source_rank, - node_world_size=node_world_size, + current_rank=current_rank, ) - for placement in placements_by_layer + for rank_to_logic_expert_ids in rank_to_logic_expert_ids_by_layer ] - assert maps_by_layer.shape == (3, 8, 4) - assert counts_by_layer.shape == (3, 8) - assert maps_by_layer.dtype == counts_by_layer.dtype == torch.int32 - assert torch.equal(maps_by_layer, torch.stack([item[0] for item in expected_by_layer])) - assert torch.equal(counts_by_layer, torch.stack([item[1] for item in expected_by_layer])) - if source_rank is None: - # expert 0 的主副本在 rank 0,且 rank 1、2 都有一个冗余副本;第 1、2 - # 个冗余副本必须分别写入映射表的第 1、2 列,而不能互相覆盖。 - assert torch.equal(counts_by_layer[:, 0], torch.tensor([3, 3, 3], dtype=torch.int32)) - assert torch.equal( - maps_by_layer[:, 0, :3], - torch.tensor([[0, 6, 10], [0, 6, 10], [0, 6, 10]], dtype=torch.int32), - ) - positions = torch.arange(maps_by_layer.shape[-1]).view(1, 1, -1) - valid = positions < counts_by_layer.unsqueeze(-1) - assert torch.all(maps_by_layer[valid] >= 0) - assert torch.all(maps_by_layer[~valid] == -1) + assert maps_by_layer == expected_by_layer + assert all( + physical_expert_id >= 0 + for logical_map in maps_by_layer + for row in logical_map + for physical_expert_id in row[2 : 2 + row[0]] + ) + assert all( + physical_expert_id == -1 + for logical_map in maps_by_layer + for row in logical_map + for physical_expert_id in row[2 + row[0] :] + ) -def test_plan_redundant_experts_prefers_local_node_load_relief(): - source_load = torch.zeros((1, 1, 2, 8), dtype=torch.int64) - source_load[0, 0, 0, 0] = 1024 +def test_plan_redundant_experts_prefers_current_rank_load_relief(): + source_load = torch.zeros((1, 1, 4, 8), dtype=torch.int64) + source_load[0, 0, 1, 0] = 1024 placement = plan_redundant_experts( source_load, num_ranks=4, num_redundant_experts_per_rank=1, expert_alignment=128, - node_world_size=2, ) assert placement[0, 1, 0] == 0 - predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) - assert torch.equal(predicted, _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128)) + predicted = _estimate_rank_load(source_load, placement, expert_alignment=128) + assert torch.equal( + predicted, + _manual_runtime_rank_load(source_load, placement, alignment=128), + ) -def test_plan_redundant_experts_single_node_matches_default_behavior(): +def test_plan_redundant_experts_accepts_aggregated_load(): load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]).unsqueeze(0).unsqueeze(2) - default = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1) - single_node = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1, node_world_size=4) - assert torch.equal(single_node, default) + placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1) + assert placement.shape == (1, 4, 1) -def test_source_node_estimate_matches_local_first_runtime_replica_sharing(): - # Expert 0 has a copy on each node. Node 0 and node 1 issue unequal - # traffic, so collapsing them before planning produces the wrong result. +def test_current_rank_estimate_matches_local_first_runtime_replica_sharing(): placement = torch.tensor([[[4], [5], [0], [1]]], dtype=torch.int64) - source_load = torch.zeros((1, 1, 2, 8), dtype=torch.int64) + source_load = torch.zeros((1, 1, 4, 8), dtype=torch.int64) source_load[0, 0, 0, 0] = 256 - source_load[0, 0, 1, 0] = 128 + source_load[0, 0, 2, 0] = 128 - predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) - runtime = _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128) - collapsed_global = _estimate_rank_load(source_load.sum(dim=2, keepdim=True), placement, expert_alignment=128) + predicted = _estimate_rank_load(source_load, placement, expert_alignment=128) + runtime = _manual_runtime_rank_load(source_load, placement, alignment=128) assert torch.equal(predicted, runtime) assert torch.equal(predicted[0, 0], torch.tensor([256.0, 0.0, 128.0, 0.0])) - assert not torch.equal(predicted, collapsed_global) -def test_source_node_planner_constraints_and_real_critical_improvement(): - source_load = torch.tensor( +def test_current_rank_planner_constraints_and_real_critical_improvement(): + load_by_node = torch.tensor( [ [ [ @@ -739,39 +773,41 @@ def test_source_node_planner_constraints_and_real_critical_improvement(): ], dtype=torch.int64, ) - initial = build_initial_redundant_expert_ids(8, 4, 1).unsqueeze(0) - planned = plan_redundant_experts(source_load, 4, 1, expert_alignment=128, node_world_size=2) + source_load = torch.zeros((3, 1, 4, 8), dtype=torch.int64) + source_load[:, :, 0] = load_by_node[:, :, 0] + source_load[:, :, 2] = load_by_node[:, :, 1] + initial = torch.tensor([[[2], [4], [6], [0]]], dtype=torch.int64) + planned = plan_redundant_experts(source_load, 4, 1, expert_alignment=128) for rank, experts in enumerate(planned[0].tolist()): assert len(experts) == len(set(experts)) == 1 assert experts[0] // 2 != rank - before = _estimate_rank_load(source_load, initial, expert_alignment=128, node_world_size=2) - after = _estimate_rank_load(source_load, planned, expert_alignment=128, node_world_size=2) - manual_before = _manual_runtime_rank_load(source_load, initial, 2, 128) - manual_after = _manual_runtime_rank_load(source_load, planned, 2, 128) + before = _estimate_rank_load(source_load, initial, expert_alignment=128) + after = _estimate_rank_load(source_load, planned, expert_alignment=128) + manual_before = _manual_runtime_rank_load(source_load, initial, 128) + manual_after = _manual_runtime_rank_load(source_load, planned, 128) assert torch.equal(before, manual_before) assert torch.equal(after, manual_after) assert after.max(dim=2).values.sum() < before.max(dim=2).values.sum() -def test_source_node_select_uses_the_same_runtime_critical_prediction(): - source_load = torch.zeros((2, 1, 2, 8), dtype=torch.int64) - source_load[:, 0, 0, 0] = torch.tensor([1024, 768]) - source_load[:, 0, 1, 6] = torch.tensor([896, 1024]) +def test_current_rank_select_uses_the_same_runtime_critical_prediction(): + source_load = torch.zeros((2, 1, 4, 8), dtype=torch.int64) + source_load[:, 0, 1, 0] = torch.tensor([1024, 768]) + source_load[:, 0, 0, 6] = torch.tensor([896, 1024]) current = torch.tensor([[[2], [4], [6], [0]]], dtype=torch.int64) - candidate = plan_redundant_experts(source_load, 4, 1, expert_alignment=128, node_world_size=2) + candidate = plan_redundant_experts(source_load, 4, 1, expert_alignment=128) selected, _improved, _metrics, _before_load, _after_load = select_improving_placements( source_load, current, candidate, rebalance_gain_threshold=0.05, expert_alignment=128, - node_world_size=2, ) assert torch.equal( - _estimate_rank_load(source_load, selected, 128, 2), - _manual_runtime_rank_load(source_load, selected, 2, 128), + _estimate_rank_load(source_load, selected, 128), + _manual_runtime_rank_load(source_load, selected, 128), ) @@ -863,6 +899,16 @@ def test_align_target_placement_keeps_retained_experts_in_live_slots(): assert torch.equal(canonical, torch.tensor([[6, 5], [0, 7], [0, 1], [2, 3]])) +def test_align_target_placement_replaces_duplicate_initial_placeholders(): + current = torch.tensor([[1, 1], [3, 3]]) + target = torch.tensor([[1, 2], [3, 0]]) + + canonical = align_target_placement(current, target) + + assert torch.equal(canonical, target) + assert build_transfer_plan(current, target, num_logical_experts=4, world_size=2, node_world_size=2) + + def test_canonical_placement_keeps_transfer_rows_and_published_map_consistent(): num_logical_experts = 8 world_size = 4 @@ -885,11 +931,17 @@ def test_canonical_placement_keeps_transfer_rows_and_published_map_consistent(): for step in plan: live_rows[step.dst_rank][num_experts_per_rank + step.dst_slot] = source_rows[step.src_rank][step.src_local_row] - logical_to_physical, replica_count = build_logical_to_physical_map(canonical, num_logical_experts) - for logical_expert, count in enumerate(replica_count.tolist()): - for physical_id in logical_to_physical[logical_expert, :count].tolist(): - rank, row = divmod(physical_id, num_physical_experts_per_rank) - assert live_rows[rank][row] == logical_expert + for current_rank in range(world_size): + logical_to_physical = build_logical_to_physical_map( + _rank_to_logic_expert_ids(canonical.tolist(), num_logical_experts), + num_logical_experts, + current_rank=current_rank, + ) + for logical_expert, packed_row in enumerate(logical_to_physical): + count = packed_row[0] + for physical_id in packed_row[2 : count + 2]: + rank, row = divmod(physical_id, num_physical_experts_per_rank) + assert live_rows[rank][row] == logical_expert def test_plan_and_broadcast_publishes_canonical_placement(monkeypatch): @@ -931,7 +983,7 @@ def test_stickiness_zero_matches_unbiased_plan(): generator = torch.Generator().manual_seed(17) load = torch.randint(1, 1000, (2, 8, 16), generator=generator).unsqueeze(2) unbiased = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - unrelated = build_initial_redundant_expert_ids(16, 4, 2).unsqueeze(0).expand(8, -1, -1).clone() + unrelated = _initial_extra_expert_placement(16, 4, 2).unsqueeze(0).expand(8, -1, -1).clone() replanned = plan_redundant_experts( load, @@ -980,17 +1032,16 @@ def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_eval monkeypatch, ): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - counter = torch.tensor([[10, 20, 30, 40], [0, 0, 0, 0]], dtype=torch.int64) + counter = torch.tensor([10, 20, 30, 40], dtype=torch.int64) manager.weights = [ type( "Weight", (), { - "expert_parallel_state": _test_parallel_state( + "fuse_moe_impl": _test_moe_impl( eplb=True, route_counter=counter, recording=False, - recorded_sample_count=1, num_logical_experts=4, world_size=1, ), @@ -1011,7 +1062,7 @@ def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_eval recordings, resets, started = [], [], [] manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) + manager._reset_route_counters = lambda: resets.append(True) monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) monkeypatch.setattr( manager_module.torch.cuda, @@ -1051,7 +1102,7 @@ def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary manager._continuous_collection_end_step = None recordings, resets, started = [], [], [] manager._set_recording = lambda enabled: recordings.append((manager.prefill_steps, enabled)) - manager._reset_recorded_samples = lambda: resets.append(manager.prefill_steps) + manager._reset_route_counters = lambda: resets.append(manager.prefill_steps) monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(manager.prefill_steps)) manager._prepare_next_sampling_window() @@ -1068,17 +1119,16 @@ def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary def test_eplb_step_does_not_start_a_second_evaluation_while_one_is_pending(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - counter = torch.tensor([[10, 20, 30, 40], [0, 0, 0, 0]], dtype=torch.int64) + counter = torch.tensor([10, 20, 30, 40], dtype=torch.int64) manager.weights = [ type( "Weight", (), { - "expert_parallel_state": _test_parallel_state( + "fuse_moe_impl": _test_moe_impl( eplb=True, route_counter=counter, recording=False, - recorded_sample_count=1, num_logical_experts=4, world_size=1, ), @@ -1127,7 +1177,7 @@ def join(self): manager.step_interval = 20 manager.sampling_interval = 20 manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] recordings, logs = [], [] manager._set_recording = lambda enabled: recordings.append(enabled) monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) @@ -1168,7 +1218,7 @@ def join(self): } manager.global_rank = 1 manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] recordings, starts = [], [] manager._set_recording = lambda enabled: recordings.append(enabled) manager._start_evaluation = lambda: starts.append(True) @@ -1243,7 +1293,7 @@ def record(self, _stream): manager.step_interval = 1 manager.sampling_interval = 1 manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] manager._set_recording = lambda _enabled: None monkeypatch.setattr(manager_module.threading, "Thread", PendingThread) monkeypatch.setattr(manager_module.torch.cuda, "Event", Event) @@ -1259,10 +1309,10 @@ def record(self, _stream): assert manager._poll_evaluation() # New worker has not produced a result. -def test_manager_collects_recent_ring_samples_in_chronological_order(monkeypatch): +def test_manager_collects_aggregated_route_counters(): counters = [ - torch.tensor([[10, 11], [20, 21], [30, 31]], dtype=torch.int64), - torch.tensor([[40, 41], [50, 51], [60, 61]], dtype=torch.int64), + torch.tensor([10, 11], dtype=torch.int64), + torch.tensor([40, 41], dtype=torch.int64), ] manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.weights = [ @@ -1270,11 +1320,10 @@ def test_manager_collects_recent_ring_samples_in_chronological_order(monkeypatch "Weight", (), { - "expert_parallel_state": _test_parallel_state( + "fuse_moe_impl": _test_moe_impl( eplb=True, route_counter=counter, recording=False, - recorded_sample_count=5, num_logical_experts=2, world_size=1, ), @@ -1282,80 +1331,26 @@ def test_manager_collects_recent_ring_samples_in_chronological_order(monkeypatch )() for counter in counters ] - manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] - manager.evaluation_group = object() - metadata_sizes = [] - - def all_reduce(metadata, **_kwargs): - metadata_sizes.append(metadata.numel()) - - monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) - - samples = manager._collect_local_samples() - - assert metadata_sizes == [4] - assert torch.equal( - samples, - torch.tensor( - [ - [[30, 31], [60, 61]], - [[10, 11], [40, 41]], - [[20, 21], [50, 51]], - ], - dtype=torch.int64, - ), - ) - - -def test_manager_collects_only_two_recent_sparse_samples(monkeypatch): - counters = [ - torch.tensor([[10, 11], [20, 21], [30, 31], [40, 41]], dtype=torch.int64), - torch.tensor([[50, 51], [60, 61], [70, 71], [80, 81]], dtype=torch.int64), - ] - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "expert_parallel_state": _test_parallel_state( - eplb=True, - route_counter=counter, - recording=False, - recorded_sample_count=2, - num_logical_experts=2, - world_size=1, - ), - }, - )() - for counter in counters - ] - manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] - manager.evaluation_group = object() - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **kwargs: None) + manager._eplb_impls = [weight.fuse_moe_impl for weight in manager.weights] + manager.num_logical_experts = 2 samples = manager._collect_local_samples() assert torch.equal( samples, - torch.tensor([[[10, 11], [50, 51]], [[20, 21], [60, 61]]], dtype=torch.int64), + torch.tensor([[[10, 11], [40, 41]]], dtype=torch.int64), ) -def test_eplb_counter_capacity_covers_default_dense_interval(monkeypatch): +def test_eplb_route_counter_has_one_entry_per_logical_expert(monkeypatch): args = type( "Args", (), {"eplb_num_redundant_experts_per_rank": 2}, )() - weight = object.__new__(fused_weight_module.FusedMoeWeight) - weight.n_routed_experts = 4 - weight.global_world_size = 2 - weight.global_rank_ = 0 - weight.enable_ep_moe = True - monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) - monkeypatch.setattr(fused_weight_module, "get_prefill_eplb_step_interval", lambda: 20) - monkeypatch.setattr(fused_weight_module, "get_node_world_size", lambda: 2) + monkeypatch.setattr(deepgemm_module, "get_env_start_args", lambda: args) + monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) original_zeros = torch.zeros @@ -1363,34 +1358,30 @@ def cpu_zeros(*shape, **kwargs): kwargs.pop("device", None) return original_zeros(*shape, **kwargs) - monkeypatch.setattr(fused_weight_module.torch, "zeros", cpu_zeros) + monkeypatch.setattr(deepgemm_module.torch, "zeros", cpu_zeros) - weight._init_expert_parallel_state() + impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace()) - assert weight.expert_parallel_state.eplb.route_counter.shape == (40, 4) + assert impl.route_counter.shape == (4,) -def test_steady_sampling_resets_fixed_ring_without_retained_history(): - counter = torch.ones((8, 4), dtype=torch.int64) +def test_steady_sampling_resets_aggregated_route_counter(): + counter = torch.ones((4,), dtype=torch.int64) manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - state = _test_parallel_state( + impl = _test_moe_impl( eplb=True, route_counter=counter, recording=False, - recorded_sample_count=99, num_logical_experts=4, world_size=1, - ).eplb - manager._eplb_states = [state] + ) + manager._eplb_impls = [impl] - manager._reset_recorded_samples() - manager._reset_recorded_samples() + manager._reset_route_counters() + manager._reset_route_counters() - assert state.route_counter.shape == (8, 4) - assert torch.count_nonzero(state.route_counter) == 0 - assert state.recorded_sample_count == 0 - assert not hasattr(manager, "_retained_local_samples") - assert not hasattr(manager, "_sample_history") + assert impl.route_counter.shape == (4,) + assert torch.count_nonzero(impl.route_counter) == 0 def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): @@ -1399,80 +1390,39 @@ def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): (), {"eplb_num_redundant_experts_per_rank": 0}, )() - weight = object.__new__(fused_weight_module.FusedMoeWeight) - weight.n_routed_experts = 4 - weight.global_world_size = 2 - weight.global_rank_ = 0 - weight.enable_ep_moe = True - monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) - - weight._init_expert_parallel_state() - - assert weight.expert_parallel_state is not None - assert weight.expert_parallel_state.eplb is None - assert weight.expert_parallel_state.num_total_physical_experts == weight.n_routed_experts - assert weight._initial_redundant_expert_ids == [] - assert not hasattr(weight, "route_counter") - assert not hasattr(weight, "routed_expert_counter_tensor") - - -def test_disable_eplb_model_init_skips_eplb_state(monkeypatch): - args = type( - "Args", - (), - {"eplb_num_redundant_experts_per_rank": 2}, - )() - weight = object.__new__(fused_weight_module.FusedMoeWeight) - weight.n_routed_experts = 4 - weight.global_world_size = 2 - weight.global_rank_ = 0 - weight.enable_ep_moe = True - monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) - monkeypatch.setattr( - fused_weight_module, - "build_initial_redundant_expert_ids", - lambda *args, **kwargs: pytest.fail("disabled scope must not initialize EPLB"), - ) - - with disable_eplb_model_init(): - weight._init_expert_parallel_state() + monkeypatch.setattr(deepgemm_module, "get_env_start_args", lambda: args) + monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) - assert weight.expert_parallel_state.eplb is None - assert weight.expert_parallel_state.num_total_physical_experts == weight.n_routed_experts - assert weight._initial_redundant_expert_ids == [] + impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace()) + assert impl.num_redundant_experts_per_rank == 0 + assert impl.num_total_physical_experts == impl.n_routed_experts + assert impl.local_logics_expert_ids_list == [0, 1] + assert not hasattr(impl, "num_primary_experts_per_rank") + assert not hasattr(impl, "initial_local_expert_ids_by_rank") + assert not hasattr(impl, "logical_to_physical_map") + assert not hasattr(impl, "route_counter") + assert not hasattr(impl, "recording") -def test_disable_eplb_model_init_scope_restores_after_exception(): - assert not is_eplb_model_init_disabled() - with disable_eplb_model_init(): - assert is_eplb_model_init_disabled() - with disable_eplb_model_init(): - assert is_eplb_model_init_disabled() - assert is_eplb_model_init_disabled() - with pytest.raises(RuntimeError): - with disable_eplb_model_init(): - raise RuntimeError - assert not is_eplb_model_init_disabled() - - -def test_manager_evaluation_collective_preserves_source_node_axis(monkeypatch): +def test_manager_evaluation_collective_preserves_current_rank_axis(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.weights = [ type( "Weight", (), { - "expert_parallel_state": _test_parallel_state( + "fuse_moe_impl": _test_moe_impl( eplb=True, - route_counter=torch.zeros((1, 4), dtype=torch.int64), + route_counter=torch.zeros((4,), dtype=torch.int64), num_logical_experts=4, world_size=1, ) }, )() ] - manager._eplb_states = [manager.weights[0].expert_parallel_state.eplb] + manager._eplb_impls = [manager.weights[0].fuse_moe_impl] manager.global_rank = 2 manager.world_size = 4 manager.node_world_size = 2 @@ -1480,7 +1430,7 @@ def test_manager_evaluation_collective_preserves_source_node_axis(monkeypatch): manager.sampling_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 - manager.current_placement = build_initial_redundant_expert_ids(4, 4, 1).unsqueeze(0) + manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0) manager.evaluation_group = object() manager._evaluation_lock = threading.Lock() manager._evaluation_result = None @@ -1493,7 +1443,7 @@ def test_manager_evaluation_collective_preserves_source_node_axis(monkeypatch): def all_reduce(tensor, **kwargs): seen["before"] = tensor.clone() seen["group"] = kwargs["group"] - # Simulate source node 0's contribution from the other ranks. + # Simulate current rank 0's contribution from another process. tensor[:, :, 0].fill_(100) monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) @@ -1509,26 +1459,31 @@ def plan_and_broadcast(global_load): assert seen["group"] is manager.evaluation_group expected_local = local - assert seen["before"].shape == (1, 1, 2, 4) + assert seen["before"].shape == (1, 1, 4, 4) assert torch.equal(seen["before"][:, :, 0], torch.zeros_like(expected_local)) - assert torch.equal(seen["before"][:, :, 1], expected_local) + assert torch.equal(seen["before"][:, :, 1], torch.zeros_like(expected_local)) + assert torch.equal(seen["before"][:, :, 2], expected_local) + assert torch.equal(seen["before"][:, :, 3], torch.zeros_like(expected_local)) assert torch.equal(seen["global_load"][:, :, 0], torch.full_like(expected_local, 100)) - assert torch.equal(seen["global_load"][:, :, 1], expected_local) + assert torch.equal(seen["global_load"][:, :, 1], torch.zeros_like(expected_local)) + assert torch.equal(seen["global_load"][:, :, 2], expected_local) + assert torch.equal(seen["global_load"][:, :, 3], torch.zeros_like(expected_local)) assert manager._evaluation_error is None - assert manager._evaluation_result["recorded_sample_count"] == 1 assert manager._evaluation_result["sample_window_steps"] == 4 -def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_call(monkeypatch): +def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_call( + monkeypatch, +): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.weights = [ type( "Weight", (), { - "expert_parallel_state": _test_parallel_state( + "fuse_moe_impl": _test_moe_impl( eplb=True, - route_counter=torch.zeros((1, 4), dtype=torch.int64), + route_counter=torch.zeros((4,), dtype=torch.int64), num_logical_experts=4, world_size=4, ) @@ -1536,7 +1491,7 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c )() for _ in range(3) ] - manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] + manager._eplb_impls = [weight.fuse_moe_impl for weight in manager.weights] manager.global_rank = 1 manager.world_size = 4 manager.node_world_size = 2 @@ -1544,7 +1499,7 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c manager.sampling_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 - manager.current_placement = build_initial_redundant_expert_ids(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() + manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() manager.evaluation_group = object() manager._evaluation_lock = threading.Lock() manager._evaluation_result = None @@ -1578,7 +1533,8 @@ def prepare_transfer(layer_plans): original_build_maps_for_layers = manager_module.build_logical_to_physical_maps_for_layers def build_maps_for_layers(*args, **kwargs): - calls.append(args[0].shape) + placements = args[0] + calls.append((len(placements), len(placements[0]), len(placements[0][0]))) return original_build_maps_for_layers(*args, **kwargs) monkeypatch.setattr( @@ -1590,7 +1546,7 @@ def build_maps_for_layers(*args, **kwargs): manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) assert manager._evaluation_error is None - assert calls == [torch.Size([2, 4, 1])] + assert calls == [(2, 4, 2)] metadata = manager._evaluation_result["metadata"] assert metadata[1] is None assert [layer_index for layer_index, _plan in manager._evaluation_result["layer_plans"]] == [0, 2] @@ -1599,13 +1555,11 @@ def build_maps_for_layers(*args, **kwargs): for layer_index in (0, 2): item = metadata[layer_index] expected = build_logical_to_physical_map( - planned_placement[layer_index], + _rank_to_logic_expert_ids(planned_placement[layer_index].tolist(), 4), 4, - source_rank=manager.global_rank, - node_world_size=manager.node_world_size, + current_rank=manager.global_rank, ) - assert torch.equal(item[0], expected[0]) - assert torch.equal(item[1], expected[1]) + assert torch.equal(item, torch.tensor(expected, dtype=torch.int32)) def test_manager_preparation_error_is_saved_as_evaluation_error(monkeypatch): @@ -1615,23 +1569,23 @@ def test_manager_preparation_error_is_saved_as_evaluation_error(monkeypatch): "Weight", (), { - "expert_parallel_state": _test_parallel_state( + "fuse_moe_impl": _test_moe_impl( eplb=True, - route_counter=torch.zeros((1, 2), dtype=torch.int64), + route_counter=torch.zeros((2,), dtype=torch.int64), num_logical_experts=2, world_size=1, ) }, )() ] - manager._eplb_states = [manager.weights[0].expert_parallel_state.eplb] + manager._eplb_impls = [manager.weights[0].fuse_moe_impl] manager.global_rank = 0 manager.world_size = 1 manager.node_world_size = 1 manager.step_interval = 20 manager.sampling_interval = 20 manager.num_logical_experts = 2 - manager.current_placement = torch.tensor([[[0], [1]]], dtype=torch.int64) + manager.current_placement = torch.tensor([[[0]]], dtype=torch.int64) manager.evaluation_group = object() manager._evaluation_lock = threading.Lock() manager._evaluation_result = None @@ -1640,7 +1594,7 @@ def test_manager_preparation_error_is_saved_as_evaluation_error(monkeypatch): manager._collect_local_samples = lambda: torch.ones((1, 1, 2), dtype=torch.int64) manager._plan_and_broadcast = lambda _global_load: { "kind": "planned", - "placement": torch.tensor([[[0], [1]]], dtype=torch.int64), + "placement": torch.tensor([[[0]]], dtype=torch.int64), "improved": torch.tensor([True]), } @@ -1668,7 +1622,7 @@ def low_latency_dispatch(self, **kwargs): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) impl.quant_method = type("Quant", (), {"method_name": "fp8"})() impl.n_routed_experts = 128 - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + _set_deepgemm_runtime(impl, _test_moe_impl(eplb=True)) logical_ids = torch.tensor([[0, 127]], dtype=torch.int32) physical_ids = torch.tensor([[128, 143]], dtype=torch.int32) impl._select_experts = lambda **_kwargs: ( @@ -1711,7 +1665,7 @@ def test_select_returns_logical_ids_and_applies_expert_scale(monkeypatch): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) impl.routed_scaling_factor = 2.0 - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + _set_deepgemm_runtime(impl, _test_moe_impl(eplb=True)) logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) calls = [] @@ -1742,7 +1696,7 @@ def test_eplb_prefill_repairs_ids_after_selection(monkeypatch): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) impl.routed_scaling_factor = 1.0 impl.quant_method = object() - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + _set_deepgemm_runtime(impl, _test_moe_impl(eplb=True)) logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) physical_ids = torch.tensor([[130, 131]], dtype=torch.long) calls = [] @@ -1773,7 +1727,7 @@ def repair(**kwargs): assert topk_idx.dtype is torch.long assert qinput == "qinput" assert calls[0]["logical_topk_ids"] is logical_ids - assert not calls[0]["record_load"] + assert not calls[0]["update_logical_expert_counter"] def test_eplb_prefill_dispatch_consumes_physical_ids_and_event(monkeypatch): @@ -1791,12 +1745,12 @@ def dispatch(self, _qinput, **kwargs): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) impl.routed_scaling_factor = 1.0 impl.quant_method = object() - state = _test_parallel_state( + runtime = _test_moe_impl( eplb=True, - route_counter=torch.zeros((3, 128), dtype=torch.int64), + route_counter=torch.zeros((128,), dtype=torch.int64), recording=True, ) - _set_expert_parallel_state(impl, state) + _set_deepgemm_runtime(impl, runtime) impl.ep_balance_counters = None calls, repair_calls = [], [] logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) @@ -1840,9 +1794,7 @@ def repair(**kwargs): assert topk_idx is physical_ids assert len(repair_calls) == 1 assert repair_calls[0]["logical_topk_ids"] is logical_ids - assert repair_calls[0]["sample_index"] == 0 - assert repair_calls[0]["record_load"] - assert state.eplb.recorded_sample_count == 1 + assert repair_calls[0]["update_logical_expert_counter"] assert calls[0]["topk_idx"] is physical_ids assert calls[0]["topk_idx"].dtype is torch.long assert calls[0]["previous_event"] is caller_event @@ -1861,7 +1813,7 @@ def dispatch(self, _qinput, **kwargs): ) impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + _set_deepgemm_runtime(impl, _test_moe_impl(eplb=True)) impl.ep_balance_counters = None calls = [] caller_event = object() @@ -1884,16 +1836,38 @@ def dispatch(self, _qinput, **kwargs): assert calls[0]["topk_idx"].dtype is torch.long -def test_deepgemm_constructor_configures_eplb(): - state = _validated_expert_parallel_state() - impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace(), expert_parallel_state=state) - assert impl.expert_parallel_state is state +def test_deepgemm_constructor_owns_eplb_runtime(monkeypatch): + monkeypatch.setattr( + deepgemm_module, + "get_env_start_args", + lambda: SimpleNamespace(eplb_num_redundant_experts_per_rank=1), + ) + monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) + original_zeros = torch.zeros + + def cpu_zeros(*shape, **kwargs): + kwargs.pop("device", None) + return original_zeros(*shape, **kwargs) + + monkeypatch.setattr(deepgemm_module.torch, "zeros", cpu_zeros) + impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace()) + + assert impl.num_primary_experts_per_rank == 2 + assert impl.num_redundant_experts_per_rank == 1 + assert impl.num_total_physical_experts == 6 + assert impl.route_counter.shape == (4,) + assert impl.recording + assert impl.local_logics_expert_ids_list == [0, 1, 1] + assert not hasattr(impl, "initial_local_expert_ids_by_rank") + assert not hasattr(impl, "expert_parallel_state") def test_eplb_prepare_repairs_logical_ids(monkeypatch): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - state = _test_parallel_state(eplb=True, recording=True) - _set_expert_parallel_state(impl, state) + runtime = _test_moe_impl(eplb=True, recording=True) + _set_deepgemm_runtime(impl, runtime) logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) physical_ids = torch.tensor([[13, 14]], dtype=torch.int32) calls = [] @@ -1908,14 +1882,14 @@ def repair(**kwargs): assert weights.tolist() == [[1.0, 1.0]] assert selected is physical_ids assert calls[0]["logical_topk_ids"] is logical_ids - assert calls[0]["record_load"] + assert calls[0]["update_logical_expert_counter"] def test_decode_masked_group_gemm_uses_all_physical_rows_when_eplb_is_enabled( monkeypatch, ): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True, num_logical_experts=8, world_size=1)) + _set_deepgemm_runtime(impl, _test_moe_impl(eplb=True, num_logical_experts=8, world_size=1)) captured = {} def masked(*args, **kwargs): @@ -1942,9 +1916,14 @@ def test_decode_fused_experts_uses_full_weight_packs_and_physical_experts( ): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) impl.n_routed_experts = 128 - _set_expert_parallel_state( + _set_deepgemm_runtime( impl, - _test_parallel_state(eplb=True, num_logical_experts=128, world_size=16, num_redundant_experts_per_rank=2), + _test_moe_impl( + eplb=True, + num_logical_experts=128, + world_size=16, + num_redundant_experts_per_rank=2, + ), ) impl.quant_method = object() impl.ep_balance_counters = None @@ -2385,7 +2364,11 @@ def copy_batch(batch, _prepared_batch): plans = [(0, []), (1, []), (2, [])] prepared_batches = [(batch, None) for batch in transfer._make_batches(plans)] - monkeypatch.setattr(transfer, "_make_batches", lambda _plans: pytest.fail("start must reuse prepared batches")) + monkeypatch.setattr( + transfer, + "_make_batches", + lambda _plans: pytest.fail("start must reuse prepared batches"), + ) transfer.start(plans, prepared_batches) deadline = time.monotonic() + 2 while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: @@ -2467,7 +2450,7 @@ def test_manager_rearms_after_rebalance_for_interval_one(): manager._steady_collection_end_step = None manager._continuous_collection_start_step = 0 manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] manager.target_placement = torch.zeros(1) manager.in_flight_started_at = 0 manager.global_rank = 1 @@ -2499,12 +2482,12 @@ def join(self): manager.sampling_interval = 20 manager.prefill_steps = 37 manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] manager._continuous_collection_start_step = None manager._continuous_collection_end_step = None recordings, resets, logs = [], [], [] manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) + manager._reset_route_counters = lambda: resets.append(True) monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) assert not manager._poll_evaluation() @@ -2556,12 +2539,12 @@ def join(self): manager.sampling_interval = 20 manager.prefill_steps = 60 manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] manager._continuous_collection_start_step = 40 manager._continuous_collection_end_step = 60 recordings, resets = [], [] manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) + manager._reset_route_counters = lambda: resets.append(True) assert not manager._poll_evaluation() assert not hasattr(manager, "_retained_local_samples") @@ -2582,7 +2565,7 @@ def test_begin_continuous_collection_uses_full_window_at_fixed_boundary(monkeypa manager._steady_collection_end_step = manager.prefill_steps + 1 recordings, resets, starts = [], [], [] manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) + manager._reset_route_counters = lambda: resets.append(True) monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) # The window is never truncated to the next boundary: it waits until 40, @@ -2614,7 +2597,7 @@ def test_begin_continuous_collection_preserves_full_window_at_sparse_boundary(): manager._steady_collection_end_step = None recordings = [] manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: None + manager._reset_route_counters = lambda: None manager._begin_continuous_collection() assert manager._continuous_collection_start_step == 140 @@ -2644,7 +2627,7 @@ def join(self): manager.sampling_interval = 20 manager._continuous_collection_start_step = 0 manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] recordings = [] manager._set_recording = lambda enabled: recordings.append(enabled) @@ -2689,7 +2672,7 @@ def join(self): manager.step_interval = 20 manager.sampling_interval = 20 manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] for expected_interval in (80, 320, 320): manager.evaluation_in_flight = True @@ -2718,7 +2701,7 @@ def test_sparse_backoff_arms_and_evaluates_only_at_new_interval_boundary(monkeyp manager.evaluation_in_flight = False recordings, resets, starts = [], [], [] manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) + manager._reset_route_counters = lambda: resets.append(True) monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) manager.step() @@ -2761,7 +2744,7 @@ def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): (), {"start": lambda self, plans, prepared_batches: setattr(self, "started", (plans, prepared_batches))}, )() - manager._reset_recorded_samples = lambda: None + manager._reset_route_counters = lambda: None prepared_batches = [object()] manager._start_rebalance( @@ -2803,7 +2786,7 @@ def test_first_rebalance_completion_switches_to_four_step_sparse_window(monkeypa manager.sampling_interval = manager.step_interval manager._steady_collection_end_step = None manager.weights = [] - manager._eplb_states = [] + manager._eplb_impls = [] manager.target_placement = torch.zeros(1) manager.in_flight_started_at = 0 manager.global_rank = 1 @@ -2911,8 +2894,20 @@ def enqueue(self, descriptor, stream): ] ] transfer._push_staging_row_layout = { - 1: [[("w13.weight", 4000, 32), ("w13.weight_scale", 5000, 32), ("w2.weight", 6000, 32)]], - 2: [[("w13.weight", 7000, 32), ("w13.weight_scale", 8000, 32), ("w2.weight", 9000, 32)]], + 1: [ + [ + ("w13.weight", 4000, 32), + ("w13.weight_scale", 5000, 32), + ("w2.weight", 6000, 32), + ] + ], + 2: [ + [ + ("w13.weight", 7000, 32), + ("w13.weight_scale", 8000, 32), + ("w2.weight", 9000, 32), + ] + ], } transfer._get_remote_read = lambda *_args: None transfer._wait_xfers = lambda _xfers: None @@ -2924,8 +2919,16 @@ def enqueue(self, descriptor, stream): staging = object() batch = [(0, run_a + run_b + [local_inbound, remote_inbound], 0, staging)] prepared = transfer._prepare_batch(batch) - monkeypatch.setattr(transfer, "_prepare_batch", lambda _batch: pytest.fail("hot path must not prepare descriptors")) - monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) + monkeypatch.setattr( + transfer, + "_prepare_batch", + lambda _batch: pytest.fail("hot path must not prepare descriptors"), + ) + monkeypatch.setattr( + transfer_module.torch.cuda, + "stream", + lambda _stream: pytest.fail("must not switch streams"), + ) transfer._copy_batch(batch, prepared) expected = ( @@ -3070,7 +3073,9 @@ def unavailable(): raise failure monkeypatch.setattr( - transfer_module._EPLBTransferBase, "__init__", lambda self, *_args: setattr(self, "device", "mock") + transfer_module._EPLBTransferBase, + "__init__", + lambda self, *_args: setattr(self, "device", "mock"), ) monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", unavailable) @@ -3392,7 +3397,11 @@ def synchronize(self): ) transfer._wait_xfers = waited_xfers.extend - monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) + monkeypatch.setattr( + transfer_module.torch.cuda, + "stream", + lambda _stream: pytest.fail("must not switch streams"), + ) prepared_push = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") transfer._copy_batch([], prepared_push) assert enqueued == [("push", stream.cuda_stream)] @@ -3418,7 +3427,11 @@ def synchronize(self): enqueued = [] transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) transfer._wait_xfers = lambda _xfers: None - monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) + monkeypatch.setattr( + transfer_module.torch.cuda, + "stream", + lambda _stream: pytest.fail("must not switch streams"), + ) prepared = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") transfer._copy_batch([], prepared) @@ -3433,12 +3446,12 @@ def test_manager_constructs_nixl_transfer(monkeypatch): (), { "n_routed_experts": 4, - "expert_parallel_state": _test_parallel_state( + "fuse_moe_impl": _test_moe_impl( eplb=True, num_logical_experts=4, world_size=2, num_redundant_experts_per_rank=2, - route_counter=torch.zeros((2, 4), dtype=torch.int64), + route_counter=torch.zeros((4,), dtype=torch.int64), ), }, )() @@ -3480,47 +3493,52 @@ def new_group(*args, **kwargs): assert "rebalance_gain_threshold=0.0700" in logs[0] assert manager._continuous_collection_start_step is None assert manager._continuous_collection_end_step == manager.step_interval - assert weight.expert_parallel_state.eplb.recording - assert manager._eplb_states[0] is weight.expert_parallel_state.eplb - assert not hasattr(weight.expert_parallel_state.eplb, "record_load") + assert weight.fuse_moe_impl.recording + assert manager._eplb_impls[0] is weight.fuse_moe_impl + assert not hasattr(weight.fuse_moe_impl, "update_logical_expert_counter") @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") -@pytest.mark.parametrize("record_load", [False, True]) +@pytest.mark.parametrize("update_logical_expert_counter", [False, True]) @pytest.mark.parametrize("tokens", [1, 32]) -def test_eplb_repair_topk_ids_maps_and_counts(record_load, tokens): - from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import eplb_repair_topk_ids +def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tokens): + from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( + eplb_repair_topk_ids, + ) topk = 4 experts = 64 logical_ids = (torch.arange(tokens * topk, dtype=torch.int32, device="cuda") % experts).view(tokens, topk) original_logical_ids = logical_ids.clone() + logical_experts = torch.arange(experts, dtype=torch.int32, device="cuda") + replica_counts = torch.where( + logical_experts % 3 == 0, + torch.full_like(logical_experts, 2), + torch.ones_like(logical_experts), + ) + has_local_replica = (logical_experts % 5 == 0).to(torch.int32) logical_to_physical = torch.stack( ( - torch.arange(experts, dtype=torch.int32, device="cuda"), - torch.arange(experts, dtype=torch.int32, device="cuda") + experts, + replica_counts, + has_local_replica, + logical_experts, + torch.where(replica_counts == 2, logical_experts + experts, logical_experts), ), dim=1, ) - logical_replica_count = torch.where( - torch.arange(experts, device="cuda") % 3 == 0, - torch.full((experts,), 2, dtype=torch.int32, device="cuda"), - torch.ones((experts,), dtype=torch.int32, device="cuda"), - ) - counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") + counter = torch.zeros((experts,), dtype=torch.int64, device="cuda") expected_counter = torch.zeros_like(counter) - if tokens == 1: - replica_indices = torch.zeros_like(logical_ids) - else: - token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) - replica_indices = ( - (((token_indices * 2654435769) & 0xFFFFFFFF) + ((logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF)) - & 0xFFFFFFFF - ) % logical_replica_count[logical_ids.to(torch.long)].to(torch.int64) - expected_ids = logical_to_physical[logical_ids.to(torch.long), replica_indices.to(torch.long)] - if record_load: - expected_counter[1].scatter_add_( + logical_ids_long = logical_ids.to(torch.long) + token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) + replica_indices = ( + (((token_indices * 2654435769) & 0xFFFFFFFF) + ((logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF)) + & 0xFFFFFFFF + ) % replica_counts[logical_ids_long].to(torch.int64) + replica_indices = torch.where(has_local_replica[logical_ids_long] != 0, 0, replica_indices) + expected_ids = logical_to_physical[logical_ids_long, replica_indices + 2] + if update_logical_expert_counter: + expected_counter.scatter_add_( 0, logical_ids.reshape(-1).to(torch.long), torch.ones(logical_ids.numel(), dtype=torch.int64, device="cuda"), @@ -3529,10 +3547,8 @@ def test_eplb_repair_topk_ids_maps_and_counts(record_load, tokens): physical_ids = eplb_repair_topk_ids( logical_topk_ids=logical_ids, logical_to_physical_map=logical_to_physical, - logical_replica_count=logical_replica_count, - expert_counter=counter, - sample_index=1, - record_load=record_load, + logical_expert_counter=counter, + update_logical_expert_counter=update_logical_expert_counter, ) torch.cuda.synchronize() @@ -3543,18 +3559,26 @@ def test_eplb_repair_topk_ids_maps_and_counts(record_load, tokens): @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") def test_eplb_repair_topk_ids_empty_input_skips_kernel(): - from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import eplb_repair_topk_ids + from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( + eplb_repair_topk_ids, + ) experts = 64 logical_ids = torch.empty((0, 4), dtype=torch.int32, device="cuda") - counter = torch.zeros((1, experts), dtype=torch.int64, device="cuda") + counter = torch.zeros((experts,), dtype=torch.int64, device="cuda") + logical_to_physical = torch.stack( + ( + torch.ones((experts,), dtype=torch.int32, device="cuda"), + torch.ones((experts,), dtype=torch.int32, device="cuda"), + torch.arange(experts, dtype=torch.int32, device="cuda"), + ), + dim=1, + ) physical_ids = eplb_repair_topk_ids( logical_topk_ids=logical_ids, - logical_to_physical_map=torch.zeros((experts, 1), dtype=torch.int32, device="cuda"), - logical_replica_count=torch.ones((experts,), dtype=torch.int32, device="cuda"), - expert_counter=counter, - sample_index=0, - record_load=True, + logical_to_physical_map=logical_to_physical, + logical_expert_counter=counter, + update_logical_expert_counter=True, ) assert physical_ids.shape == (0, 4) diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 3eff70f940..b334a6fa15 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -1,9 +1,11 @@ """NIXL EPLB correctness tests and a two-GPU 512 MiB micro-performance test.""" + import os import random import socket import statistics import time +from types import SimpleNamespace import pytest import torch @@ -16,16 +18,23 @@ build_transfer_plan, ) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_initial_redundant_expert_ids, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - EPLBState, - ExpertParallelState, + build_initial_local_expert_ids, ) pytest.importorskip("nixl", reason="NIXL package is required") +def _initial_extra_expert_placement(num_logical_experts, world_size, num_redundant_experts_per_rank): + num_primary_experts_per_rank = num_logical_experts // world_size + initial_local_expert_ids_by_rank = build_initial_local_expert_ids( + num_logical_experts, world_size, num_redundant_experts_per_rank + ) + return torch.tensor( + [expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids_by_rank], + dtype=torch.int64, + ) + + class _Pack: def __init__(self, weight, weight_scale): self.weight = weight @@ -44,16 +53,18 @@ def _free_port(): class _FakeWeight: def __init__(self, rank, layer_index, row_elements): self.n_routed_experts = 32 - self.expert_parallel_state = ExpertParallelState( - num_logical_experts=32, - world_size=2, - eplb=EPLBState( - num_redundant_experts_per_rank=16, - initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids(32, 2, 16), - logical_to_physical_map=torch.zeros((32, 2), dtype=torch.int32, device="cuda"), - logical_replica_count=torch.ones(32, dtype=torch.int32, device="cuda"), - route_counter=torch.zeros((1, 32), dtype=torch.int64, device="cuda"), + self.fuse_moe_impl = SimpleNamespace( + n_routed_experts=32, + num_primary_experts_per_rank=16, + num_redundant_experts_per_rank=16, + logical_to_physical_map=torch.cat( + ( + torch.ones((32, 1), dtype=torch.int32, device="cuda"), + torch.zeros((32, 3), dtype=torch.int32, device="cuda"), + ), + dim=1, ), + route_counter=torch.zeros((32,), dtype=torch.int64, device="cuda"), ) base = rank * 100 + layer_index * 100 self.w13 = self._pack(base, row_elements) @@ -183,7 +194,7 @@ def _eplb_worker(rank, port, queue): weights = [_FakeWeight(rank, layer_index, row_elements) for layer_index in range(layer_count)] transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=2) assert transfer.staging_depth == 8 - assert transfer._eplb_states[0] is weights[0].expert_parallel_state.eplb + assert transfer._eplb_impls[0] is weights[0].fuse_moe_impl assert all(tensor.is_cuda for staging in transfer.staging for _, tensor in staging) wrap_layer_plans = [(layer_index, plan) for layer_index in range(layer_count)] @@ -235,16 +246,18 @@ def _depth_value(expert, layer_index, offset): class _DepthWeight: def __init__(self, rank, layer_index, initial_placement): self.n_routed_experts = 256 - self.expert_parallel_state = ExpertParallelState( - num_logical_experts=256, - world_size=8, - eplb=EPLBState( - num_redundant_experts_per_rank=4, - initial_redundant_expert_ids_by_rank=initial_placement.clone(), - logical_to_physical_map=torch.zeros((256, 8), dtype=torch.int32, device="cuda"), - logical_replica_count=torch.ones(256, dtype=torch.int32, device="cuda"), - route_counter=torch.zeros((1, 256), dtype=torch.int64, device="cuda"), + self.fuse_moe_impl = SimpleNamespace( + n_routed_experts=256, + num_primary_experts_per_rank=32, + num_redundant_experts_per_rank=4, + logical_to_physical_map=torch.cat( + ( + torch.ones((256, 1), dtype=torch.int32, device="cuda"), + torch.zeros((256, 9), dtype=torch.int32, device="cuda"), + ), + dim=1, ), + route_counter=torch.zeros((256,), dtype=torch.int64, device="cuda"), ) logical_ids = list(range(rank * 32, (rank + 1) * 32)) + initial_placement[rank].tolist() self.w13 = self._pack(logical_ids, layer_index, 0) @@ -354,7 +367,7 @@ def _depth_worker(rank, port): dist.init_process_group("gloo", rank=rank, world_size=8) control_group = dist.new_group(list(range(8)), backend="gloo") transfer_group = dist.new_group(list(range(8)), backend="gloo") - initial_placement = build_initial_redundant_expert_ids(256, 8, 4) + initial_placement = _initial_extra_expert_placement(256, 8, 4) weights = [_DepthWeight(rank, layer_index, initial_placement) for layer_index in range(9)] transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=8) assert transfer.staging_depth == 8 @@ -367,7 +380,10 @@ def _depth_worker(rank, port): layer_index: align_target_placement(current, _depth_target(layer_index)) for layer_index in range(9) } first_plans = [ - (layer_index, build_transfer_plan(current, first_targets[layer_index], 256, 8, 8)) + ( + layer_index, + build_transfer_plan(current, first_targets[layer_index], 256, 8, 8), + ) for layer_index in first_order ] _assert_peer_coverage(first_plans) From 11c81b4a7bfabbdcea876afe3478559f03a067ea Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 12:14:08 +0000 Subject: [PATCH 18/72] style: minimize DeepGEMM formatting diff --- .../fused_moe/impl/deepgemm_impl.py | 24 ++++--------------- 1 file changed, 4 insertions(+), 20 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 f6ef8e8e0a..1b2c4fced4 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 @@ -23,9 +23,7 @@ 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.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( eplb_repair_topk_ids, ) @@ -88,9 +86,7 @@ def _select_experts( per_expert_scale: Optional[torch.Tensor] = None, ): """Select logical experts without applying the EPLB physical layout.""" - from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import ( - select_experts, - ) + from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts topk_weights, topk_ids = select_experts( hidden_states=input_tensor, @@ -257,14 +253,7 @@ def hook(): compute_load=compute_load, ) - return ( - recv_x, - recv_topk_idx, - recv_topk_weights, - handle.num_recv_tokens_per_expert_list, - handle, - hook, - ) + return recv_x, recv_topk_idx, recv_topk_weights, handle.num_recv_tokens_per_expert_list, handle, hook def masked_group_gemm( self, @@ -347,12 +336,7 @@ def low_latency_combine( handle: Any, ): combined_x, event_overlap, hook = dist_group_manager.ep_low_latency_buffer.low_latency_combine( - gemm_out_b, - topk_idx, - topk_weights, - handle, - async_finish=False, - return_recv_hook=True, + gemm_out_b, topk_idx, topk_weights, handle, async_finish=False, return_recv_hook=True ) return combined_x, hook From 1265836bf63bb115bc4fe1372bdcaeb89ff43cf0 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 13:26:20 +0000 Subject: [PATCH 19/72] refactor(eplb): simplify expert balance metrics --- lightllm/common/basemodel/basemodel.py | 3 - .../meta_weights/fused_moe/ep_balance.py | 14 - .../fused_moe/impl/deepgemm_impl.py | 14 - .../fused_moe/grouped_fused_moe_ep.py | 10 - lightllm/distributed/communication_op.py | 9 - lightllm/server/api_cli.py | 5 - lightllm/server/core/objs/start_args_type.py | 1 - lightllm/server/metrics/metrics.py | 15 +- .../mode_backend/ep_balance_monitor.py | 346 ----------- .../model_infer/mode_backend/eplb_manager.py | 28 + .../server/router/model_infer/model_rpc.py | 9 - unit_tests/common/fused_moe/test_eplb.py | 37 +- .../model_infer/test_ep_balance_monitor.py | 546 ------------------ 13 files changed, 66 insertions(+), 971 deletions(-) delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py delete mode 100644 lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py delete mode 100644 unit_tests/server/router/model_infer/test_ep_balance_monitor.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index ea9d95927c..5fb8d60d0b 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -71,7 +71,6 @@ class TpPartBaseModel: def __init__(self, kvargs): self.args = get_env_start_args() self.eplb_manager = None - self.ep_balance_monitor = None self.run_mode = kvargs["run_mode"] self.weight_dir_ = kvargs["weight_dir"] self.max_total_token_num = kvargs["max_total_token_num"] @@ -324,8 +323,6 @@ def forward(self, model_input: ModelInput): return self._decode(model_input) def _after_prefill(self): - if self.ep_balance_monitor is not None: - self.ep_balance_monitor.record_prefill_round() if self.eplb_manager is not None: self.eplb_manager.step() diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py deleted file mode 100644 index 70435146c5..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py +++ /dev/null @@ -1,14 +0,0 @@ -from dataclasses import dataclass - - -@dataclass(slots=True) -class PrefillEPBalanceCounters: - """Cumulative CPU loads for one EP MoE layer's completed prefill dispatches.""" - - route_load: int = 0 - compute_load: int = 0 - - def accumulate(self, route_load: int, compute_load: int): - """Accumulate exact route and alignment-expanded compute loads for one prefill dispatch.""" - self.route_load += route_load - self.compute_load += compute_load 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 1b2c4fced4..e19bb63634 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 @@ -34,7 +34,6 @@ class FuseMoeDeepGEMM(FuseMoeBaseImpl): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self._init_eplb_runtime() - self.ep_balance_counters = None def _init_eplb_runtime(self): world_size = get_global_world_size() @@ -141,7 +140,6 @@ def _fused_experts( quant_method=self.quant_method, is_prefill=is_prefill, previous_event=None, # for overlap - ep_balance_counters=self.ep_balance_counters, ) return output @@ -238,20 +236,8 @@ def dispatch( use_tma_aligned_col_major_sf=True, ) - counters = self.ep_balance_counters - route_load = compute_load = 0 - if counters is not None: - # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. - route_load = topk_idx.numel() - compute_load = recv_x[0].shape[0] - def hook(): event.current_stream_wait() - if counters is not None: - counters.accumulate( - route_load=route_load, - compute_load=compute_load, - ) return recv_x, recv_topk_idx, recv_topk_weights, handle.num_recv_tokens_per_expert_list, handle, hook 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 ff430e5e51..ca39376bab 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 @@ -19,7 +19,6 @@ ep_gather_chunk, ep_zero_padding, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters from lightllm.utils.envs_utils import ( get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, @@ -202,7 +201,6 @@ def fused_experts( quant_method: Any, is_prefill: Optional[bool], previous_event: Optional[Any] = None, - ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): check_ep_expert_dtype(quant_method) if use_sm100_mega_moe(quant_method): @@ -224,7 +222,6 @@ def fused_experts( w1_scale=w13.weight_scale, w2_scale=w2.weight_scale, previous_event=previous_event, - ep_balance_counters=ep_balance_counters, ) @@ -243,7 +240,6 @@ def fused_experts_impl( w1_scale: Optional[torch.Tensor] = None, w2_scale: Optional[torch.Tensor] = None, previous_event: Optional[Any] = None, - ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): # Check constraints. assert hidden_states.shape[1] == w1.shape[2], "Hidden size mismatch" @@ -294,12 +290,6 @@ def fused_experts_impl( do_expand=True, use_tma_aligned_col_major_sf=True, ) - if ep_balance_counters is not None: - # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. - ep_balance_counters.accumulate( - route_load=topk_idx.numel(), - compute_load=recv_x[0].shape[0], - ) # 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 diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index a9ceef84f9..375df2d6f2 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -107,7 +107,6 @@ def all_gather_into_tensor(self, output_: torch.Tensor, input_: torch.Tensor, as class DistributeGroupManager: def __init__(self): self.groups = [] - self.ep_balance_monitor_group = None self.ep_buffer = None self.ep_low_latency_buffer = None self.ep_mega_moe_buffer = None @@ -125,14 +124,6 @@ def create_groups(self, group_size: int): if not args.disable_flashinfer_allreduce: group.init_flashinfer_reduce() self.groups.append(group) - if ( - getattr(args, "enable_ep_moe", False) - and not getattr(args, "disable_ep_balance_monitor", False) - and getattr(args, "run_mode", "normal") != "decode" - and not getattr(args, "enable_prefill_cudagraph", False) - and not is_sm100_gpu() - ): - self.ep_balance_monitor_group = dist.new_group(ranks=list(range(get_global_world_size())), backend="gloo") return def get_default_group(self) -> CustomProcessGroup: diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index d615edcc72..23848c20bd 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -766,11 +766,6 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="""Whether to enable ep moe for deepseekv3 model.""", ) - parser.add_argument( - "--disable_ep_balance_monitor", - action="store_true", - help="""Disable the prefill expert balance monitor enabled by default for EP-MoE.""", - ) parser.add_argument( "--eplb_num_redundant_experts_per_rank", type=int, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index bec1d435c4..4527dd6387 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -186,7 +186,6 @@ class StartArgs: default="gpu_counter", metadata={"choices": ["cpu_counter", "pin_mem_counter", "gpu_counter"]} ) enable_ep_moe: bool = field(default=False) - disable_ep_balance_monitor: bool = field(default=False) eplb_num_redundant_experts_per_rank: int = field(default=0) enable_fused_shared_experts: bool = field(default=False) mtp_mode: Optional[str] = field( diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index c19d756c17..9b8c620322 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -32,14 +32,9 @@ "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", "lightllm_gen_throughput": "Generation throughput of latest completed request (tokens/s)", "lightllm_num_running_reqs": "Number of running requests", - "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token": ( - "Estimated critical-path excess compute per logical routed source token, GFLOPs/token" - ), - "lightllm_prefill_ep_compute_critical_overhead_ratio": ( - "Estimated excess critical compute divided by balanced compute; 0.3 means +30%" - ), - "lightllm_prefill_ep_placement_pressure_drift": ( - "Normalized temporal drift of overloaded-rank pressure from the latest complete prefill report" + "lightllm_eplb_topk_expert_imbalance_ratio": ( + "Maximum routed token count divided by the mean across logical experts, averaged across MoE layers in the " + "latest EPLB sample window" ), } @@ -120,9 +115,7 @@ def init_metrics(self, args): self.create_gauge("lightllm_cache_hit_rate") self.create_gauge("lightllm_gen_throughput") self.create_gauge("lightllm_num_running_reqs") - self.create_gauge("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token") - self.create_gauge("lightllm_prefill_ep_compute_critical_overhead_ratio") - self.create_gauge("lightllm_prefill_ep_placement_pressure_drift") + self.create_gauge("lightllm_eplb_topk_expert_imbalance_ratio") def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) diff --git a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py deleted file mode 100644 index 107f4ba504..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py +++ /dev/null @@ -1,346 +0,0 @@ -import threading -from array import array -from typing import Optional, Tuple - -import torch -import torch.distributed as dist - -from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters -from lightllm.distributed.communication_op import dist_group_manager -from lightllm.server.metrics.manager import MetricClient -from lightllm.utils.device_utils import is_sm100_gpu -from lightllm.utils.dist_utils import ( - get_global_rank, - get_global_world_size, -) -from lightllm.utils.log_utils import init_logger -from lightllm.utils.shm_port_args import get_shm_port_args - - -logger = init_logger(__name__) - -EP_BALANCE_PREFILL_ROUNDS_PER_REPORT = 100 -EP_BALANCE_ROUND_BUFFER_CAPACITY = 4096 -EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS = 20 -EP_BALANCE_PRESSURE_DRIFT_STABLE_THRESHOLD = 0.10 -ROUTE_LOAD = 0 -COMPUTE_LOAD = 1 -GFLOP = 1_000_000_000 - - -def should_enable_ep_balance_monitor(args) -> bool: - if args.enable_prefill_cudagraph or is_sm100_gpu(): - return False - return args.enable_ep_moe and not args.disable_ep_balance_monitor and args.run_mode != "decode" - - -def calculate_prefill_balance_stats( - round_stats: torch.Tensor, # [num_rounds, num_layers, world_size, 2] (route/compute) - layer_routed_experts: torch.Tensor, # [num_layers] - layer_flops_per_expert_token: torch.Tensor, # [num_layers] - layer_topks: torch.Tensor, # [num_layers] - source_token_replication: int, - report_min_route_samples_per_expert: int = 100, -) -> Optional[dict]: - """Summarize complete-prefill samples from [round, layer, rank, route/compute]. - - MoE layers execute sequentially, and every layer waits for its slowest EP - rank. Preserve the layer dimension until after taking the cross-rank max so - that different slow ranks in different layers cannot cancel each other. - """ - assert source_token_replication > 0 - - layer_route_load = round_stats[:, :, :, ROUTE_LOAD].sum(dim=(0, 2)) - minimum_route_samples = layer_routed_experts * report_min_route_samples_per_expert - if torch.any(layer_route_load < minimum_route_samples): - return None - total_route_load = layer_route_load.sum() - - compute_rank_load = round_stats[:, :, :, COMPUTE_LOAD] - compute_rank_load_float = compute_rank_load.to(torch.float64) - excess_compute_load = compute_rank_load_float.max(dim=2).values - compute_rank_load_float.mean(dim=2) - - # Every MoE expert token executes the two projections packed in w13 plus - # the w2 projection. Weight each padded compute token by the layer's - # actual matrix sizes so the metric remains comparable across models. - excess_compute_flops = (excess_compute_load * layer_flops_per_expert_token.to(torch.float64)).sum() - balanced_compute_flops = ( - compute_rank_load_float.mean(dim=2) * layer_flops_per_expert_token.to(torch.float64) - ).sum() - if balanced_compute_flops == 0: - return None - - # Non-TPSP prefill gathers one route-load copy per TP rank. Divide the - # replica count out so GFLOP/token uses logical source tokens. - source_tokens = total_route_load.to(torch.float64) / ( - layer_topks.to(torch.float64).sum() * source_token_replication - ) - if source_tokens == 0: - return None - - return { - "prefill_rounds": int(compute_rank_load.shape[0]), - "critical_overhead_gflops_per_routed_token": float((excess_compute_flops / source_tokens / GFLOP).item()), - "prefill_ep_compute_critical_overhead_ratio": float((excess_compute_flops / balanced_compute_flops).item()), - } - - -def calculate_prefill_placement_pressure_drift( - round_stats: torch.Tensor, - previous_pressure_signature: Optional[torch.Tensor] = None, - bucket_rounds: int = EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS, -) -> Tuple[float, torch.Tensor]: - """Measure how prefill rank-pressure placement changes across time buckets. - - This is rank-0 CPU-only report analysis. The returned final bucket is a - compact signature that allows the next report to include the boundary pair. - """ - if bucket_rounds <= 0: - raise ValueError(f"bucket_rounds must be positive, got {bucket_rounds}") - if round_stats.ndim != 4 or round_stats.shape[-1] != 2: - raise ValueError( - "round_stats must have shape [num_rounds, num_layers, world_size, 2], " f"got {tuple(round_stats.shape)}" - ) - num_rounds, num_layers, world_size, _ = round_stats.shape - if num_rounds <= 0: - raise ValueError("round_stats must contain at least one round") - if num_layers <= 0 or world_size <= 0: - raise ValueError("round_stats must contain at least one layer and rank") - if num_rounds % bucket_rounds != 0: - raise ValueError(f"num_rounds ({num_rounds}) must be divisible by bucket_rounds ({bucket_rounds})") - - num_buckets = num_rounds // bucket_rounds - bucket_rank_load = ( - round_stats[:, :, :, COMPUTE_LOAD] - .to(torch.float64) - .reshape(num_buckets, bucket_rounds, num_layers, world_size) - .sum(dim=1) - ) - mean_rank_load = bucket_rank_load.mean(dim=2, keepdim=True).clamp_min(1) - pressure = torch.relu(bucket_rank_load / mean_rank_load - 1) - - if previous_pressure_signature is not None: - expected_shape = (num_layers, world_size) - if tuple(previous_pressure_signature.shape) != expected_shape: - raise ValueError( - "previous_pressure_signature must have shape " - f"{expected_shape}, got {tuple(previous_pressure_signature.shape)}" - ) - left = torch.cat((previous_pressure_signature.to(torch.float64).unsqueeze(0), pressure[:-1]), dim=0) - right = pressure - else: - left = pressure[:-1] - right = pressure[1:] - - total_pressure = (left + right).sum() - if total_pressure == 0: - drift = 0.0 - else: - drift = float((left - right).abs().sum().div(total_pressure).item()) - return drift, pressure[-1].clone() - - -def classify_prefill_placement_pressure_drift(drift: float) -> str: - if drift < EP_BALANCE_PRESSURE_DRIFT_STABLE_THRESHOLD: - return "stable" - return "dynamic" - - -def _find_fused_moe_weights(model): - weights_by_id = {} - for layer in model.trans_layers_weight: - for value in getattr(layer, "__dict__", {}).values(): - if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: - weights_by_id[id(value)] = value - return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) - - -class EPBalanceMonitor: - """Report cross-rank imbalance for non-overlapping blocks of complete prefill rounds.""" - - def __init__(self, model: TpPartBaseModel): - self.global_rank = get_global_rank() - self.world_size = get_global_world_size() - self.weights = _find_fused_moe_weights(model) - self.enabled = bool(self.weights) - - if not self.enabled: - return - - self.source_token_replication = 1 if model.args.enable_tpsp_mix_mode else model.tp_world_size_ - self.counters: list[PrefillEPBalanceCounters] = [PrefillEPBalanceCounters() for _ in self.weights] - for weight, counter in zip(self.weights, self.counters): - weight.fuse_moe_impl.ep_balance_counters = counter - self.layer_routed_experts = torch.tensor( - [weight.n_routed_experts for weight in self.weights], dtype=torch.int64 - ) - self.layer_flops_per_expert_token = torch.tensor( - [ - # Each expert-token performs gate, up, and down projections; each MAC counts as 2 FLOPs. - 2 * 3 * weight.hidden_size * weight.moe_intermediate_size - for weight in self.weights - ], - dtype=torch.float64, - ) - self.layer_topks = torch.tensor( - [weight.num_experts_per_tok for weight in self.weights], - dtype=torch.float64, - ) - self._round_buffer_storage = array("q", [0]) * (EP_BALANCE_ROUND_BUFFER_CAPACITY * len(self.weights) * 2) - self._round_buffer = torch.frombuffer(self._round_buffer_storage, dtype=torch.int64).view( - EP_BALANCE_ROUND_BUFFER_CAPACITY, len(self.weights), 2 - ) - self._round_ready = threading.Event() - self._written_round_count = 0 # Prefill rounds fully written to the ring buffer. - self._processed_round_count = 0 # Prefill rounds consumed by the monitor thread. - self._overflowed = False - self._previous_pressure_signature: Optional[torch.Tensor] = None - self._common_round_end = torch.zeros((), dtype=torch.int64) - - self.gloo_group = dist_group_manager.ep_balance_monitor_group - if self.gloo_group is None: - raise RuntimeError("EP balance monitor requires a pre-created dedicated Gloo process group") - self.metric_client = MetricClient(get_shm_port_args().metric_port) if self.global_rank == 0 else None - threading.Thread(target=self._monitor_loop, daemon=True, name="ep-balance-monitor").start() - - def record_prefill_round(self): - """Publish one complete all-layer prefill sample to the SPSC ring.""" - if not self.enabled: - return - - written_round_count = self._written_round_count - if written_round_count - self._processed_round_count >= EP_BALANCE_ROUND_BUFFER_CAPACITY: - if not self._overflowed: - self._overflowed = True - self._round_ready.set() - return - - storage_index = (written_round_count % EP_BALANCE_ROUND_BUFFER_CAPACITY) * len(self.counters) * 2 - for counter in self.counters: - self._round_buffer_storage[storage_index] = counter.route_load - self._round_buffer_storage[storage_index + 1] = counter.compute_load - counter.route_load = 0 - counter.compute_load = 0 - storage_index += 2 - - # Publish only after the entire slot is written. The SPSC producer and - # monitor thread run under the CPython GIL, so this count is the release - # point for the corresponding ring slot. - self._written_round_count = written_round_count + 1 - if self._written_round_count - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: - self._round_ready.set() - - def _get_common_round_end(self) -> int: - """Return the exclusive round boundary completed by every rank.""" - self._common_round_end.fill_(self._written_round_count) - dist.all_reduce(self._common_round_end, op=dist.ReduceOp.MIN, group=self.gloo_group) - return int(self._common_round_end.item()) - - def _raise_buffer_overflow(self, phase: str, common_round_end: Optional[int] = None): - message = ( - "EP balance prefill-round buffer overflowed " - f"phase={phase} written={self._written_round_count} " - f"processed={self._processed_round_count} capacity={EP_BALANCE_ROUND_BUFFER_CAPACITY}" - ) - if common_round_end is not None: - message += f" common_round_end={common_round_end}" - raise RuntimeError(message) - - def _copy_local_rounds(self, start: int, end: int) -> torch.Tensor: - """Copy local prefill-round loads in the half-open range [start, end).""" - num_rounds = end - start - if num_rounds > EP_BALANCE_ROUND_BUFFER_CAPACITY: - raise ValueError("requested EP balance round range exceeds ring capacity") - start_index = start % EP_BALANCE_ROUND_BUFFER_CAPACITY - if start_index + num_rounds <= EP_BALANCE_ROUND_BUFFER_CAPACITY: - return self._round_buffer[start_index : start_index + num_rounds].clone() - end_index = (start_index + num_rounds) % EP_BALANCE_ROUND_BUFFER_CAPACITY - return torch.cat((self._round_buffer[start_index:], self._round_buffer[:end_index]), dim=0) - - def _gather_round_stats(self, local_round_stats: torch.Tensor) -> Optional[torch.Tensor]: - """Gather rank-local stats as [round, layer, rank, route/compute].""" - gathered = ( - [torch.empty_like(local_round_stats) for _ in range(self.world_size)] if self.global_rank == 0 else None - ) - dist.gather(local_round_stats, gather_list=gathered, dst=0, group=self.gloo_group) - if self.global_rank != 0: - return None - # [rank, round, layer, route/compute] - # -> [round, layer, rank, route/compute] - return torch.stack(gathered).permute(1, 2, 0, 3) - - def _log_stats(self, round_stats: torch.Tensor): - """Compute and log balance statistics for one complete global window.""" - compute = calculate_prefill_balance_stats( - round_stats, - self.layer_routed_experts, - self.layer_flops_per_expert_token, - self.layer_topks, - self.source_token_replication, - ) - if compute is None: - return - - drift, self._previous_pressure_signature = calculate_prefill_placement_pressure_drift( - round_stats, - previous_pressure_signature=self._previous_pressure_signature, - ) - drift_state = classify_prefill_placement_pressure_drift(drift) - - logger.info( - "ep_balance " - f"phase=prefill prefill_rounds={compute['prefill_rounds']} " - "prefill_ep_critical_overhead_gflops_per_routed_token=" - f"{compute['critical_overhead_gflops_per_routed_token']:.4f} " - "prefill_ep_compute_critical_overhead_ratio=" - f"{compute['prefill_ep_compute_critical_overhead_ratio']:.4f} " - f"prefill_ep_placement_pressure_drift={drift:.4f} " - f"prefill_ep_placement_pressure_state={drift_state}" - ) - self.metric_client.gauge_set( - "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", - compute["critical_overhead_gflops_per_routed_token"], - ) - self.metric_client.gauge_set( - "lightllm_prefill_ep_compute_critical_overhead_ratio", - compute["prefill_ep_compute_critical_overhead_ratio"], - ) - self.metric_client.gauge_set("lightllm_prefill_ep_placement_pressure_drift", drift) - - def _monitor_loop(self): - """Consume commonly completed rounds in background report-sized windows.""" - try: - while True: - self._round_ready.wait() - self._round_ready.clear() - if self._overflowed: - self._raise_buffer_overflow("before_sync") - common_round_end = self._get_common_round_end() - if common_round_end - self._processed_round_count > EP_BALANCE_ROUND_BUFFER_CAPACITY: - self._raise_buffer_overflow("common_round_lag", common_round_end=common_round_end) - - while common_round_end - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: - round_start = self._processed_round_count - round_end = round_start + EP_BALANCE_PREFILL_ROUNDS_PER_REPORT - local_round_stats = self._copy_local_rounds(round_start, round_end) - if self._overflowed: - self._raise_buffer_overflow("after_copy") - round_stats = self._gather_round_stats(local_round_stats) - self._processed_round_count = round_end - if self.global_rank == 0: - self._log_stats(round_stats) - - if self._written_round_count - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: - self._round_ready.set() - except Exception as exc: - logger.exception(f"EP balance monitor stopped unexpectedly: {exc}") - self._disable() - return - - def _disable(self): - """Detach counters from MoE weights and disable monitoring.""" - for weight in self.weights: - weight.fuse_moe_impl.ep_balance_counters = None - self.enabled = False diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index e411a6abdb..dccba5d99e 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -20,6 +20,7 @@ align_target_placement, build_transfer_plan, ) +from lightllm.server.metrics.manager import MetricClient from lightllm.utils.dist_utils import ( get_global_rank, get_global_world_size, @@ -31,12 +32,14 @@ get_prefill_eplb_step_interval, ) from lightllm.utils.log_utils import init_logger +from lightllm.utils.shm_port_args import get_shm_port_args logger = init_logger(__name__) EPLB_MIN_AVG_TOKENS_PER_EXPERT = 100 EPLB_EXPERT_ALIGNMENT = 128 EPLB_CONTROL_ERROR = -1 EPLB_STEADY_SAMPLE_STEPS = 4 +EPLB_EXPERT_IMBALANCE_RATIO_METRIC = "lightllm_eplb_topk_expert_imbalance_ratio" class EPLBManager: @@ -81,6 +84,7 @@ def __init__(self, model: TpPartBaseModel): self._evaluation_result = None self._evaluation_error = None self._evaluation_thread = None + self.metric_client = None # A fresh manager starts with one continuous base window. After a # sufficient evaluation, steady state returns to the cheap sparse # probe. An insufficient sparse probe schedules one fresh continuous @@ -200,6 +204,13 @@ def _collect_local_samples(self) -> torch.Tensor: raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") return torch.stack(counters).unsqueeze(0).cpu() + def _publish_expert_load_metrics(self, result): + if self.global_rank != 0 or "expert_imbalance_ratio" not in result: + return + if self.metric_client is None: + self.metric_client = MetricClient(get_shm_port_args().metric_port) + self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) + def _commit_layer_metadata(self, layer_index: int): impl = self._eplb_impls[layer_index] impl.logical_to_physical_map.copy_(self.target_metadata[layer_index], non_blocking=True) @@ -334,6 +345,7 @@ def _evaluate_after_event(self, event: torch.cuda.Event): global_load[:, :, self.global_rank] = local_load dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) result = self._plan_and_broadcast(global_load) + result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) result["sample_window_steps"] = sample_window_steps if result["kind"] == "planned": metadata = [None] * len(self.weights) @@ -417,6 +429,7 @@ def _poll_evaluation(self): self._evaluation_thread.join() self.evaluation_in_flight = False self._evaluation_thread = None + self._publish_expert_load_metrics(result) if result["kind"] == "insufficient": from_continuous_window = self._continuous_collection_end_step is not None if from_continuous_window: @@ -537,6 +550,21 @@ def _imbalance_summary(rank_load: torch.Tensor) -> Dict[str, float]: } +def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: + """Average each layer's maximum-to-mean logical-expert token ratio.""" + if global_load.ndim != 4: + raise ValueError("global_load must be [samples, layers, ranks, logical_experts]") + if global_load.shape[1] == 0 or global_load.shape[3] == 0: + raise ValueError("global_load must contain at least one layer and logical expert") + layer_expert_load = global_load.sum(dim=(0, 2)).to(torch.float64) + layer_means = layer_expert_load.mean(dim=1) + valid_layers = layer_means > 0 + if not torch.any(valid_layers): + return 0.0 + layer_ratios = layer_expert_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] + return float(layer_ratios.mean().item()) + + def _find_fused_moe_weights(model): weights_by_id = {} for layer in model.trans_layers_weight: diff --git a/lightllm/server/router/model_infer/model_rpc.py b/lightllm/server/router/model_infer/model_rpc.py index e4784e79bd..29b8b857a3 100644 --- a/lightllm/server/router/model_infer/model_rpc.py +++ b/lightllm/server/router/model_infer/model_rpc.py @@ -26,10 +26,6 @@ PDDPForDecodeNode, ) from lightllm.server.router.model_infer.mode_backend.rl_backend_ops import RlBackendOps -from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( - EPBalanceMonitor, - should_enable_ep_balance_monitor, -) from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.utils.log_utils import init_logger from lightllm.utils.graceful_utils import graceful_registry @@ -100,11 +96,6 @@ def exposed_init_model(self, kvargs): logger.info(f"use {self.backend.__class__.__name__}") self.backend.init_model(kvargs) self.rl_backend_ops = RlBackendOps(self.backend) if self.args.enable_rl else None - - if should_enable_ep_balance_monitor(self.args): - monitor = EPBalanceMonitor(self.backend.model) - if monitor.enabled: - self.backend.model.ep_balance_monitor = monitor return def exposed_get_max_total_token_num(self): diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 7ca506f767..da4324e382 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1342,6 +1342,39 @@ def test_manager_collects_aggregated_route_counters(): ) +def test_expert_load_imbalance_ratio_averages_layer_ratios(): + global_load = torch.tensor( + [ + [ + [[1, 2, 3], [1, 2, 3]], + [[4, 5, 6], [6, 5, 4]], + ] + ], + dtype=torch.int64, + ) + + ratio = manager_module._expert_load_imbalance_ratio(global_load) + + assert ratio == pytest.approx(1.25) + + +def test_manager_publishes_expert_load_metrics_from_rank_zero(): + calls = [] + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.global_rank = 0 + manager.metric_client = SimpleNamespace(gauge_set=lambda name, value: calls.append((name, value))) + + manager._publish_expert_load_metrics( + { + "expert_imbalance_ratio": 1.25, + } + ) + + assert calls == [ + (manager_module.EPLB_EXPERT_IMBALANCE_RATIO_METRIC, 1.25), + ] + + def test_eplb_route_counter_has_one_entry_per_logical_expert(monkeypatch): args = type( "Args", @@ -1469,6 +1502,7 @@ def plan_and_broadcast(global_load): assert torch.equal(seen["global_load"][:, :, 2], expected_local) assert torch.equal(seen["global_load"][:, :, 3], torch.zeros_like(expected_local)) assert manager._evaluation_error is None + assert manager._evaluation_result["expert_imbalance_ratio"] == 1.0 assert manager._evaluation_result["sample_window_steps"] == 4 @@ -1751,7 +1785,6 @@ def dispatch(self, _qinput, **kwargs): recording=True, ) _set_deepgemm_runtime(impl, runtime) - impl.ep_balance_counters = None calls, repair_calls = [], [] logical_ids = torch.tensor([[3, 4]], dtype=torch.int32) physical_ids = torch.tensor([[130, 131]], dtype=torch.long) @@ -1814,7 +1847,6 @@ def dispatch(self, _qinput, **kwargs): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) _set_deepgemm_runtime(impl, _test_moe_impl(eplb=True)) - impl.ep_balance_counters = None calls = [] caller_event = object() monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_buffer", Buffer()) @@ -1926,7 +1958,6 @@ def test_decode_fused_experts_uses_full_weight_packs_and_physical_experts( ), ) impl.quant_method = object() - impl.ep_balance_counters = None captured = [] def fused(**kwargs): diff --git a/unit_tests/server/router/model_infer/test_ep_balance_monitor.py b/unit_tests/server/router/model_infer/test_ep_balance_monitor.py deleted file mode 100644 index 2ddcaafddf..0000000000 --- a/unit_tests/server/router/model_infer/test_ep_balance_monitor.py +++ /dev/null @@ -1,546 +0,0 @@ -import threading -from array import array -from types import SimpleNamespace - -import pytest -import torch -from prometheus_client import generate_latest - -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters -from lightllm.distributed import communication_op as communication_op_module -from lightllm.server.metrics.metrics import Monitor -from lightllm.server.router.model_infer.mode_backend import ep_balance_monitor as monitor_module -from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( - EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS, - calculate_prefill_placement_pressure_drift, - calculate_prefill_balance_stats, - classify_prefill_placement_pressure_drift, - should_enable_ep_balance_monitor, -) - - -@pytest.fixture(autouse=True) -def _mock_non_sm100(monkeypatch): - monkeypatch.setattr(monitor_module, "is_sm100_gpu", lambda: False) - - -def _stats(source_token_replication: int): - return calculate_prefill_balance_stats( - torch.tensor([[[[800, 40], [800, 20]], [[800, 20], [800, 40]]]], dtype=torch.int64), - layer_routed_experts=torch.tensor([1, 1]), - layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), - layer_topks=torch.tensor([2.0, 2.0]), - source_token_replication=source_token_replication, - ) - - -def _monitor_args(**overrides): - args = { - "enable_ep_moe": True, - "disable_ep_balance_monitor": False, - "run_mode": "normal", - "enable_prefill_cudagraph": False, - } - args.update(overrides) - return SimpleNamespace(**args) - - -def _pressure_round_stats(bucket_rank_loads): - """Build [round, layer=1, rank, route/compute] CPU samples for drift tests.""" - return torch.tensor([[[[0, load] for load in rank_loads]] for rank_loads in bucket_rank_loads], dtype=torch.int64) - - -def test_pressure_drift_is_zero_for_identical_pressure(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [2, 1]]), bucket_rounds=1) - assert drift == 0.0 - - -def test_pressure_drift_is_one_for_complete_hot_rank_migration(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 0], [0, 2]]), bucket_rounds=1) - assert drift == 1.0 - - -def test_pressure_drift_tracks_same_rank_magnitude_change(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [3, 1]]), bucket_rounds=1) - assert drift == pytest.approx(0.2) - - -def test_pressure_drift_is_invariant_to_uniform_load_scale(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [4, 2]]), bucket_rounds=1) - assert drift == 0.0 - - -def test_pressure_drift_is_zero_for_balanced_inputs(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[8, 8], [16, 16]]), bucket_rounds=1) - assert drift == 0.0 - - -def test_pressure_drift_previous_signature_bridges_report_boundary(): - _, signature = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 0]]), bucket_rounds=1) - drift, next_signature = calculate_prefill_placement_pressure_drift( - _pressure_round_stats([[0, 2]]), previous_pressure_signature=signature, bucket_rounds=1 - ) - assert drift == 1.0 - assert torch.equal(next_signature, torch.tensor([[0.0, 1.0]], dtype=torch.float64)) - - -def test_pressure_drift_default_bucket_handles_normal_report_window(): - round_stats = _pressure_round_stats([[4, 2]] * 100) - drift, signature = calculate_prefill_placement_pressure_drift(round_stats) - assert EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS == 20 - assert drift == 0.0 - assert signature.shape == (1, 2) - - -@pytest.mark.parametrize( - ("drift", "expected"), - [ - (0.0, "stable"), - (0.0999, "stable"), - (0.10, "dynamic"), - (0.2999, "dynamic"), - (0.30, "dynamic"), - (0.75, "dynamic"), - (1.0, "dynamic"), - ], -) -def test_pressure_drift_classification_boundaries(drift, expected): - assert classify_prefill_placement_pressure_drift(drift) == expected - - -def test_monitor_log_stats_reports_pressure_drift_and_bridges_reports(monkeypatch): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.layer_routed_experts = torch.tensor([1], dtype=torch.int64) - monitor.layer_flops_per_expert_token = torch.tensor([2.0], dtype=torch.float64) - monitor.layer_topks = torch.tensor([2.0], dtype=torch.float64) - monitor.source_token_replication = 1 - monitor._previous_pressure_signature = None - metric_calls = [] - monitor.metric_client = SimpleNamespace(gauge_set=lambda name, value: metric_calls.append((name, value))) - logs = [] - monkeypatch.setattr(monitor_module.logger, "info", logs.append) - - def report_for_hot_rank(hot_rank): - round_stats = torch.zeros((100, 1, 2, 2), dtype=torch.int64) - round_stats[:, :, :, monitor_module.ROUTE_LOAD] = 100 - round_stats[:, :, hot_rank, monitor_module.COMPUTE_LOAD] = 2 - monitor._log_stats(round_stats) - - report_for_hot_rank(0) - first_signature = monitor._previous_pressure_signature.clone() - report_for_hot_rank(1) - - assert "prefill_ep_placement_pressure_drift=0.0000" in logs[0] - assert "prefill_ep_placement_pressure_state=stable" in logs[0] - assert "prefill_ep_placement_pressure_drift=0.2000" in logs[1] - assert "prefill_ep_placement_pressure_state=dynamic" in logs[1] - assert torch.equal(first_signature, torch.tensor([[1.0, 0.0]], dtype=torch.float64)) - assert torch.equal(monitor._previous_pressure_signature, torch.tensor([[0.0, 1.0]], dtype=torch.float64)) - assert metric_calls == [ - ("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", pytest.approx(2e-11)), - ("lightllm_prefill_ep_compute_critical_overhead_ratio", pytest.approx(1.0)), - ("lightllm_prefill_ep_placement_pressure_drift", 0.0), - ("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", pytest.approx(2e-11)), - ("lightllm_prefill_ep_compute_critical_overhead_ratio", pytest.approx(1.0)), - ("lightllm_prefill_ep_placement_pressure_drift", pytest.approx(0.2)), - ] - - -def test_critical_overhead_preserves_per_layer_slowest_rank(): - stats = _stats(source_token_replication=1) - assert stats is not None - assert stats["critical_overhead_gflops_per_routed_token"] == pytest.approx(7.5e-11) - assert stats["prefill_ep_compute_critical_overhead_ratio"] == pytest.approx(1 / 3) - - -def test_non_tpsp_tp_replication_scales_gflops_per_token_but_not_ratio(): - tpsp_stats = _stats(source_token_replication=1) - non_tpsp_tp8_stats = _stats(source_token_replication=8) - assert tpsp_stats is not None and non_tpsp_tp8_stats is not None - assert non_tpsp_tp8_stats["critical_overhead_gflops_per_routed_token"] == pytest.approx( - tpsp_stats["critical_overhead_gflops_per_routed_token"] * 8 - ) - assert non_tpsp_tp8_stats["prefill_ep_compute_critical_overhead_ratio"] == pytest.approx( - tpsp_stats["prefill_ep_compute_critical_overhead_ratio"] - ) - - -def test_critical_overhead_is_zero_when_ranks_are_balanced(): - stats = calculate_prefill_balance_stats( - torch.tensor([[[[100, 32], [100, 32]], [[100, 64], [100, 64]]]], dtype=torch.int64), - layer_routed_experts=torch.tensor([1, 1]), - layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), - layer_topks=torch.tensor([2.0, 2.0]), - source_token_replication=1, - ) - assert stats is not None - assert stats["critical_overhead_gflops_per_routed_token"] == 0.0 - assert stats["prefill_ep_compute_critical_overhead_ratio"] == 0.0 - - -@pytest.mark.parametrize( - "round_stats", - [ - torch.tensor([[[[1, 1], [1, 1]]]], dtype=torch.int64), - torch.tensor([[[[100, 0], [100, 0]]]], dtype=torch.int64), - ], -) -def test_critical_overhead_rejects_insufficient_or_zero_compute_samples(round_stats): - assert ( - calculate_prefill_balance_stats( - round_stats, - layer_routed_experts=torch.tensor([1]), - layer_flops_per_expert_token=torch.tensor([2.0]), - layer_topks=torch.tensor([2.0]), - source_token_replication=1, - ) - is None - ) - - -def test_cpu_counter_accumulates_multiple_prefill_dispatches(): - counters = PrefillEPBalanceCounters() - counters.accumulate(route_load=3, compute_load=128) - counters.accumulate(route_load=4, compute_load=256) - assert (counters.route_load, counters.compute_load) == (7, 384) - - -def test_monitor_reuses_manager_precreated_dedicated_gloo_group(monkeypatch): - sentinel_group = object() - impl = SimpleNamespace(ep_balance_counters=None) - weight = SimpleNamespace( - fuse_moe_impl=impl, - n_routed_experts=8, - hidden_size=16, - moe_intermediate_size=32, - num_experts_per_tok=2, - ) - model = SimpleNamespace(args=SimpleNamespace(enable_tpsp_mix_mode=True), tp_world_size_=1) - - monkeypatch.setattr(monitor_module, "_find_fused_moe_weights", lambda _: [weight]) - monkeypatch.setattr(monitor_module, "get_global_rank", lambda: 0) - monkeypatch.setattr(monitor_module, "get_global_world_size", lambda: 2) - metric_ports = [] - monkeypatch.setattr(monitor_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=4321)) - monkeypatch.setattr(monitor_module, "MetricClient", lambda port: metric_ports.append(port) or object()) - monkeypatch.setattr(monitor_module, "dist_group_manager", SimpleNamespace(ep_balance_monitor_group=sentinel_group)) - monkeypatch.setattr(monitor_module.dist, "new_group", lambda *args, **kwargs: pytest.fail("unexpected new_group")) - monkeypatch.setattr( - monitor_module.threading, - "Thread", - lambda *args, **kwargs: SimpleNamespace(start=lambda: None), - ) - - monitor = monitor_module.EPBalanceMonitor(model) - - assert monitor.gloo_group is sentinel_group - assert impl.ep_balance_counters is monitor.counters[0] - assert metric_ports == [4321] - - -def test_nonzero_rank_monitor_does_not_create_metric_client(monkeypatch): - sentinel_group = object() - impl = SimpleNamespace(ep_balance_counters=None) - weight = SimpleNamespace( - fuse_moe_impl=impl, - n_routed_experts=8, - hidden_size=16, - moe_intermediate_size=32, - num_experts_per_tok=2, - ) - model = SimpleNamespace(args=SimpleNamespace(enable_tpsp_mix_mode=True), tp_world_size_=1) - metric_client_calls = [] - monkeypatch.setattr(monitor_module, "_find_fused_moe_weights", lambda _: [weight]) - monkeypatch.setattr(monitor_module, "get_global_rank", lambda: 1) - monkeypatch.setattr(monitor_module, "get_global_world_size", lambda: 2) - monkeypatch.setattr(monitor_module, "dist_group_manager", SimpleNamespace(ep_balance_monitor_group=sentinel_group)) - monkeypatch.setattr(monitor_module.dist, "new_group", lambda *args, **kwargs: pytest.fail("unexpected new_group")) - monkeypatch.setattr(monitor_module, "get_shm_port_args", lambda: pytest.fail("unexpected port lookup")) - monkeypatch.setattr( - monitor_module, - "MetricClient", - lambda port: metric_client_calls.append(port) or pytest.fail("unexpected metric client"), - ) - monkeypatch.setattr( - monitor_module.threading, - "Thread", - lambda *args, **kwargs: SimpleNamespace(start=lambda: None), - ) - - monitor = monitor_module.EPBalanceMonitor(model) - - assert monitor.metric_client is None - assert metric_client_calls == [] - - -@pytest.mark.parametrize("disable_monitor", [False, True]) -def test_group_manager_creates_monitor_gloo_group_only_when_enabled(monkeypatch, disable_monitor): - monitor_group = object() - custom_groups = [] - - class FakeCustomProcessGroup: - def init_symm_mem_reduce(self): - pass - - def init_flashinfer_reduce(self): - pass - - args = SimpleNamespace( - enable_ep_moe=True, - disable_ep_balance_monitor=disable_monitor, - run_mode="normal", - enable_prefill_cudagraph=False, - disable_symm_mem_allreduce=True, - disable_flashinfer_allreduce=True, - ) - monkeypatch.setattr(communication_op_module, "get_env_start_args", lambda: args) - monkeypatch.setattr( - communication_op_module, - "CustomProcessGroup", - lambda: custom_groups.append(FakeCustomProcessGroup()) or custom_groups[-1], - ) - monkeypatch.setattr(communication_op_module, "get_global_world_size", lambda: 2) - monkeypatch.setattr(communication_op_module, "is_sm100_gpu", lambda: False) - calls = [] - monkeypatch.setattr( - communication_op_module.dist, - "new_group", - lambda *args, **kwargs: calls.append((args, kwargs)) or monitor_group, - ) - - manager = communication_op_module.DistributeGroupManager() - manager.create_groups(group_size=2) - - assert len(manager.groups) == 2 - if disable_monitor: - assert calls == [] - assert manager.ep_balance_monitor_group is None - else: - assert calls == [((), {"ranks": [0, 1], "backend": "gloo"})] - assert manager.ep_balance_monitor_group is monitor_group - - -def test_monitor_registers_prefill_ep_gauges_with_model_label(): - monitor = Monitor( - SimpleNamespace( - metric_gateway=None, - job_name="test", - grouping_key=[], - enable_monitor_auth=False, - model_name="monitor-test-model", - max_req_total_len=128, - mtp_step=0, - ) - ) - values = { - "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token": 1.25, - "lightllm_prefill_ep_compute_critical_overhead_ratio": 0.3, - "lightllm_prefill_ep_placement_pressure_drift": 0.125, - } - assert set(values).issubset(monitor.monitor_registry) - for name, value in values.items(): - monitor.gauge_set(name, value) - - exposition = generate_latest(monitor.registry).decode() - for name, value in values.items(): - assert f'{name}{{model_name="monitor-test-model"}} {value}' in exposition - - -def test_record_prefill_round_stores_cumulative_counter_deltas_in_ring_buffer(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.enabled = True - monitor.counters = [PrefillEPBalanceCounters()] - monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) - monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( - monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 - ) - monitor._round_ready = threading.Event() - monitor._written_round_count = 0 - monitor._processed_round_count = 0 - monitor._overflowed = False - - monitor.counters[0].accumulate(route_load=3, compute_load=128) - monitor.record_prefill_round() - assert (monitor.counters[0].route_load, monitor.counters[0].compute_load) == (0, 0) - monitor.counters[0].accumulate(route_load=2, compute_load=256) - monitor.record_prefill_round() - - assert torch.equal( - monitor._copy_local_rounds(0, 2), - torch.tensor([[[3, 128]], [[2, 256]]], dtype=torch.int64), - ) - assert not hasattr(monitor, "_round_lock") - - -def test_spsc_ring_copy_wraps_without_a_lock(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.enabled = True - monitor.counters = [PrefillEPBalanceCounters()] - monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) - monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( - monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 - ) - monitor._round_ready = threading.Event() - monitor._written_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2 - monitor._processed_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2 - monitor._overflowed = False - - for value in (11, 12, 13): - monitor.counters[0].accumulate(route_load=value, compute_load=value * 10) - monitor.record_prefill_round() - - assert torch.equal( - monitor._copy_local_rounds(monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2, monitor._written_round_count), - torch.tensor([[[11, 110]], [[12, 120]], [[13, 130]]], dtype=torch.int64), - ) - assert not hasattr(monitor, "_round_lock") - - -def test_spsc_ring_overflow_is_deferred_to_the_monitor_thread(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.enabled = True - monitor.counters = [PrefillEPBalanceCounters(route_load=7, compute_load=70)] - monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) - monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( - monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 - ) - monitor._round_ready = threading.Event() - monitor._written_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - monitor._processed_round_count = 0 - monitor._overflowed = False - - monitor.record_prefill_round() - - assert monitor._overflowed - assert monitor._round_ready.is_set() - assert monitor._written_round_count == monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - assert (monitor.counters[0].route_load, monitor.counters[0].compute_load) == (7, 70) - - -def test_raise_buffer_overflow_always_reports_phase_and_ring_counts(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor._written_round_count = 23 - monitor._processed_round_count = 7 - - with pytest.raises(RuntimeError) as exc_info: - monitor._raise_buffer_overflow("before_sync") - - assert str(exc_info.value) == ( - "EP balance prefill-round buffer overflowed " - f"phase=before_sync written=23 processed=7 capacity={monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY}" - ) - - -def test_raise_buffer_overflow_optionally_reports_common_round_end(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor._written_round_count = 23 - monitor._processed_round_count = 7 - - with pytest.raises(RuntimeError) as exc_info: - monitor._raise_buffer_overflow("common_round_lag", common_round_end=19) - - assert str(exc_info.value) == ( - "EP balance prefill-round buffer overflowed " - f"phase=common_round_lag written=23 processed=7 capacity={monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY} " - "common_round_end=19" - ) - - -def test_gather_round_stats_only_allocates_receive_buffers_on_rank_zero(monkeypatch): - local_round_stats = torch.tensor([[[3, 128]]], dtype=torch.int64) - sentinel_group = object() - - rank_zero_monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - rank_zero_monitor.global_rank = 0 - rank_zero_monitor.world_size = 2 - rank_zero_monitor.gloo_group = sentinel_group - - def root_gather(input_tensor, gather_list, dst, group): - assert dst == 0 and group is sentinel_group - assert len(gather_list) == 2 - gather_list[0].copy_(input_tensor) - gather_list[1].copy_(input_tensor + 1) - - monkeypatch.setattr(monitor_module.dist, "gather", root_gather) - result = rank_zero_monitor._gather_round_stats(local_round_stats) - assert torch.equal(result, torch.tensor([[[[3, 128], [4, 129]]]], dtype=torch.int64)) - - nonzero_monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - nonzero_monitor.global_rank = 1 - nonzero_monitor.world_size = 2 - nonzero_monitor.gloo_group = sentinel_group - - def nonroot_gather(input_tensor, gather_list, dst, group): - assert input_tensor is local_round_stats - assert gather_list is None - assert dst == 0 and group is sentinel_group - - monkeypatch.setattr(monitor_module.dist, "gather", nonroot_gather) - assert nonzero_monitor._gather_round_stats(local_round_stats) is None - - -def test_find_fused_moe_weights_discovers_any_layer_member_once_and_sorts(monkeypatch): - class FakeFusedMoeWeight: - def __init__(self, layer_num, enabled=True): - self.layer_num_ = layer_num - self.enable_ep_moe = enabled - - monkeypatch.setattr(monitor_module, "FusedMoeWeight", FakeFusedMoeWeight) - first = FakeFusedMoeWeight(3) - second = FakeFusedMoeWeight(1) - disabled = FakeFusedMoeWeight(0, enabled=False) - model = SimpleNamespace( - trans_layers_weight=[ - SimpleNamespace(experts_=first, alias=first, ignored=disabled), - SimpleNamespace(any_direct_member=second), - ] - ) - - assert monitor_module._find_fused_moe_weights(model) == [second, first] - - -def test_monitor_disable_detaches_counters_from_all_impls(): - impl = SimpleNamespace(ep_balance_counters="unset") - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.weights = [SimpleNamespace(fuse_moe_impl=impl)] - monitor.enabled = True - monitor._disable() - assert impl.ep_balance_counters is None - assert not monitor.enabled - - -def test_critical_overhead_requires_minimum_samples_for_every_layer(): - round_stats = torch.tensor([[[[200, 32], [200, 32]], [[1, 32], [1, 32]]]], dtype=torch.int64) - assert ( - calculate_prefill_balance_stats( - round_stats, - layer_routed_experts=torch.tensor([1, 1]), - layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), - layer_topks=torch.tensor([2.0, 2.0]), - source_token_replication=1, - ) - is None - ) - - -def test_ep_moe_normal_and_prefill_enable_monitor_by_default(): - assert should_enable_ep_balance_monitor(_monitor_args()) - assert should_enable_ep_balance_monitor(_monitor_args(run_mode="prefill")) - - -def test_disable_ep_balance_monitor_turns_monitor_off(): - assert not should_enable_ep_balance_monitor(_monitor_args(disable_ep_balance_monitor=True)) - - -def test_non_ep_moe_and_decode_mode_do_not_enable_monitor(): - assert not should_enable_ep_balance_monitor(_monitor_args(enable_ep_moe=False)) - assert not should_enable_ep_balance_monitor(_monitor_args(run_mode="decode")) - - -def test_prefill_cudagraph_silently_disables_monitor(): - assert not should_enable_ep_balance_monitor(_monitor_args(run_mode="prefill", enable_prefill_cudagraph=True)) - - -def test_sm100_silently_disables_monitor(monkeypatch): - monkeypatch.setattr(monitor_module, "is_sm100_gpu", lambda: True) - assert not should_enable_ep_balance_monitor(_monitor_args()) From 7100fdf9c93d540856d3a683620ef20b9786fb7d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 17 Sep 2026 15:33:13 +0000 Subject: [PATCH 20/72] refactor(eplb): migrate expert weights through pinned memory --- .../meta_weights/fused_moe/eplb_placement.py | 18 +- .../fused_moe/fused_moe_weight.py | 14 +- lightllm/common/eplb_utils.py | 3 - .../model_infer/mode_backend/eplb_manager.py | 7 +- .../model_infer/mode_backend/eplb_transfer.py | 711 ++------------ unit_tests/common/fused_moe/test_eplb.py | 893 +++--------------- .../fused_moe/test_eplb_transfer_gpu.py | 463 ++------- 7 files changed, 293 insertions(+), 1816 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py index b43e011b8b..c5ab5b0789 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -9,28 +9,30 @@ def build_initial_local_expert_ids( ) -> list[list[int]]: """构建每个 rank 初始持有的完整 logical expert ID 列表。 - 每个 rank 先持有连续划分得到的主专家,再用本 rank 最后一个主专家 - 填充额外物理槽。这里仅负责生成 Python 列表;调用方如果要参与 tensor + 每个 rank 先持有连续划分得到的主专家,再按 rank 顺序选择不属于 + 本 rank 的专家作为默认冗余副本。这里仅负责生成 Python 列表;调用方如果要参与 tensor 运算,需要自行转换为 ``torch.Tensor``。 例如 ``num_logical_experts=8``、``num_ranks=4``、每个 rank 有 2 个 额外槽时,每个 rank 分到 2 个主专家,结果为: - ``[[0, 1, 1, 1], [2, 3, 3, 3], [4, 5, 5, 5], [6, 7, 7, 7]]`` + ``[[0, 1, 2, 3], [2, 3, 4, 5], [4, 5, 6, 7], [6, 7, 0, 1]]`` - 其中每行前两个值是主专家,后两个值是等待首次 EPLB 调整的占位副本。 + 其中每行前两个值是主专家,后两个值是已在初始加载阶段就可用的冗余副本。 """ assert num_logical_experts % num_ranks == 0 num_experts_per_rank = num_logical_experts // num_ranks - assert num_redundant_experts_per_rank >= 0 + assert 0 <= num_redundant_experts_per_rank <= num_logical_experts - num_experts_per_rank local_expert_ids_by_rank = [] for rank in range(num_ranks): first_expert_id = rank * num_experts_per_rank local_expert_ids = list(range(first_expert_id, first_expert_id + num_experts_per_rank)) - # 额外物理槽先复制本 rank 最后一个主专家。首次 EPLB 规划完成后, - # transfer 会把这些占位行替换成实际需要的跨 rank 冗余专家。 - local_expert_ids.extend([local_expert_ids[-1]] * num_redundant_experts_per_rank) + first_redundant_expert_id = ((rank + 1) * num_experts_per_rank) % num_logical_experts + local_expert_ids.extend( + (first_redundant_expert_id + offset) % num_logical_experts + for offset in range(num_redundant_experts_per_rank) + ) local_expert_ids_by_rank.append(local_expert_ids) return local_expert_ids_by_rank 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 7b34c60d0a..175b4313c7 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 @@ -90,14 +90,14 @@ def _init_weight_partition(self): self.split_inter_size = self.moe_intermediate_size // self.tp_world_size_ if self.enable_ep_moe: assert self.num_fused_shared_experts == 0, "num_fused_shared_experts must be 0 when enable_ep_moe" - self.local_expert_ids = self.fuse_moe_impl.local_logics_expert_ids_list + self.local_logic_expert_ids_list = self.fuse_moe_impl.local_logics_expert_ids_list logger.debug( f"global_rank {self.global_rank_} layerindex {self.layer_num_} " - f"local_logics_expert_ids_list: {self.local_expert_ids}" + f"local_logic_expert_ids_list: {self.local_logic_expert_ids_list}" ) - self.local_n_routed_experts = len(self.local_expert_ids) + self.local_n_routed_experts = len(self.local_logic_expert_ids_list) else: - self.local_expert_ids = list(range(self.n_routed_experts + self.num_fused_shared_experts)) + self.local_logic_expert_ids_list = list(range(self.n_routed_experts + self.num_fused_shared_experts)) def experts( self, @@ -254,7 +254,7 @@ def load_hf_weights(self, weights): # Load bias self._load_e_score_correction_bias(weights) self._load_per_expert_scale(weights) - self._load_weight(self.local_expert_ids, weights) + self._load_weight(self.local_logic_expert_ids_list, weights) def verify_load(self): weight_load_ok = all(all(_weight_pack.load_ok) for _weight_pack in self.w1_list + self.w2_list + self.w3_list) @@ -323,8 +323,8 @@ def _get_expert_weight_list(self, weight_pack: WeightPack): weight_list.append(expert_weight) return weight_list - def _load_weight(self, local_expert_ids: List[int], weights: Dict[str, torch.Tensor]): - for local_expert_idx, expert_idx in enumerate(local_expert_ids): + def _load_weight(self, local_logic_expert_ids_list: List[int], weights: Dict[str, torch.Tensor]): + for local_expert_idx, expert_idx in enumerate(local_logic_expert_ids_list): with self.lock: self._load_expert(expert_idx, local_expert_idx, weights) self._load_expert_scale( diff --git a/lightllm/common/eplb_utils.py b/lightllm/common/eplb_utils.py index 63e3dea069..4dd901a600 100644 --- a/lightllm/common/eplb_utils.py +++ b/lightllm/common/eplb_utils.py @@ -1,9 +1,6 @@ """Small, dependency-free EPLB helpers shared by transfer and model profiling.""" -EPLB_MAX_STAGING_DEPTH = 8 - - def extract_eplb_expert_tensors(weight): result = [] for pack_name in ("w13", "w2"): diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index dccba5d99e..d76c2a6106 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -16,7 +16,7 @@ select_improving_placements, ) from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( - NixlEPLBTransfer, + PinnedMemoryEPLBTransfer, align_target_placement, build_transfer_plan, ) @@ -102,7 +102,7 @@ def __init__(self, model: TpPartBaseModel): # This control-group scalar is only touched from the main inference # thread, never by the background evaluation thread. self._control_ready_count = torch.empty(1, dtype=torch.int32) - self.transfer = NixlEPLBTransfer(self.weights, self.transfer_group, self.global_rank, self.world_size) + self.transfer = PinnedMemoryEPLBTransfer(self.weights, self.transfer_group, self.global_rank, self.world_size) if self.global_rank == 0: logger.info( "eplb enabled " @@ -392,7 +392,6 @@ def _evaluate_after_event(self, event: torch.cuda.Event): ) result["metadata"] = metadata result["layer_plans"] = layer_plans - result["prepared_batches"] = self.transfer.prepare_transfer(layer_plans) with self._evaluation_lock: self._evaluation_result = result except BaseException as exc: @@ -506,7 +505,7 @@ def _start_rebalance(self, result): self.in_flight_layers = [layer_index for layer_index, _ in layer_plans] self.in_flight = True self.in_flight_started_at = time.time() - self.transfer.start(layer_plans, result["prepared_batches"]) + self.transfer.start(layer_plans) if self.global_rank == 0: actual_changed_slot_count = sum(len(plan) for _, plan in layer_plans) cross_node_transfer_count = sum( diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index eb0ac23fa3..32c6d7e9bc 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -1,17 +1,14 @@ -"""Asynchronous expert-row migration for EPLB.""" -import ctypes -import os -import re -import socket +"""Layer-by-layer expert-row migration for EPLB.""" + import threading from collections import Counter, defaultdict, deque from dataclasses import dataclass -from typing import Dict, Iterable, List, Optional, Sequence, Tuple +from typing import List, Sequence, Tuple import torch import torch.distributed as dist -from lightllm.common.eplb_utils import EPLB_MAX_STAGING_DEPTH, extract_eplb_expert_tensors +from lightllm.common.eplb_utils import extract_eplb_expert_tensors @dataclass(frozen=True) @@ -23,21 +20,12 @@ class TransferStep: def align_target_placement(current: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - """Canonicalize a target row layout without moving retained experts. - - EPLB placement is rank-based: redundant slots on one rank are - interchangeable. Retained experts therefore keep their live physical - slot, while new experts fill freed slots in the planner's target-row - order. The returned placement is the single canonical layout that must - be used both for transfers and for published routing metadata. - """ + """Canonicalize a target row layout without moving retained experts.""" assert current.ndim == target.ndim == 2 assert tuple(current.shape) == tuple(target.shape) - current_rows = current.tolist() - target_rows = target.tolist() aligned_target_rows = [] - for current_row, target_row in zip(current_rows, target_rows): + for current_row, target_row in zip(current.tolist(), target.tolist()): remaining_target = Counter(target_row) aligned_row = list(current_row) freed_slots = [] @@ -69,16 +57,10 @@ def build_transfer_plan( num_experts_per_rank = num_logical_experts // world_size current_rows = current.tolist() aligned_target_rows = align_target_placement(current, target).tolist() - # 初始占位布局允许同一个 logical expert 在一个 rank 上出现多次,因此 - # 保留所有物理来源候选,让首次迁移也可以从任意已加载的副本复制。 + # A primary row is always a valid source. Existing replicas are also + # candidates so that a destination can prefer a same-node copy. candidates_by_expert = [ - [ - ( - expert // num_experts_per_rank, - expert % num_experts_per_rank, - ) - ] - for expert in range(num_logical_experts) + [(expert // num_experts_per_rank, expert % num_experts_per_rank)] for expert in range(num_logical_experts) ] for rank, row in enumerate(current_rows): for slot, expert in enumerate(row): @@ -103,74 +85,100 @@ def build_transfer_plan( return plan -class _EPLBTransferBase: - """Shared live/staging buffers and publish/commit lifecycle.""" +class PinnedMemoryEPLBTransfer: + """Move one layer at a time through reusable pinned CPU row buffers. + + Every rank executes the same ordered Gloo broadcasts. A source first + copies a live GPU row to pinned memory; destinations then copy that row + into a single-layer GPU staging buffer. The inference thread publishes + the staging rows and routing metadata together at a safe forward boundary. + """ + + backend = "pin-memory" def __init__(self, weights, transfer_group, global_rank, world_size): self._eplb_impls = [weight.fuse_moe_impl for weight in weights] self.transfer_group = transfer_group self.global_rank = global_rank self.world_size = world_size - self.num_experts_per_rank = weights[0].fuse_moe_impl.num_primary_experts_per_rank + self.num_experts_per_rank = self._eplb_impls[0].num_primary_experts_per_rank self.device = weights[0].w13.weight.device self.live = [extract_eplb_expert_tensors(weight) for weight in weights] self._validate_live_layout() - num_redundant_slots_per_rank = self._eplb_impls[0].num_redundant_experts_per_rank + + num_redundant_slots = self._eplb_impls[0].num_redundant_experts_per_rank self.staging = [ - [ - ( - name, - torch.empty( - (num_redundant_slots_per_rank,) + tuple(tensor.shape[1:]), - dtype=tensor.dtype, - device=tensor.device, - ), - ) - for name, tensor in self.live[0] - ] - for _ in range(self.staging_depth) + ( + name, + torch.empty( + (num_redundant_slots,) + tuple(tensor.shape[1:]), + dtype=tensor.dtype, + device=tensor.device, + ), + ) + for name, tensor in self.live[0] ] - self._release = [threading.Event() for _ in range(self.staging_depth)] - for release in self._release: - release.set() - self._error = None - self._consumed_events = [torch.cuda.Event() for _ in range(self.staging_depth)] - self._consumed_recorded = [False] * self.staging_depth - self._changed_dst_slots = [()] * self.staging_depth + self.pinned_rows = [ + ( + name, + torch.empty( + tuple(tensor.shape[1:]), + dtype=tensor.dtype, + device="cpu", + pin_memory=True, + ), + ) + for name, tensor in self.live[0] + ] + self._copy_stream = torch.cuda.Stream(device=self.device) + self._release = threading.Event() + self._release.set() + self._consumed_event = torch.cuda.Event() + self._consumed_recorded = False + self._changed_dst_slots = () self._pending = deque() self._pending_lock = threading.Lock() + self._error = None self._thread = None - self._needs_staging_reuse_barrier = False def _validate_live_layout(self) -> None: reference = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in self.live[0]] - num_redundant_slots_per_rank = self._eplb_impls[0].num_redundant_experts_per_rank + num_redundant_slots = self._eplb_impls[0].num_redundant_experts_per_rank for layer_index, (impl, tensors) in enumerate(zip(self._eplb_impls, self.live)): layout = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in tensors] assert layout == reference, f"EPLB layer {layer_index} has incompatible expert tensor layout" - assert ( - impl.num_redundant_experts_per_rank == num_redundant_slots_per_rank - ), "EPLB redundant slot count must match" + assert impl.num_redundant_experts_per_rank == num_redundant_slots, "EPLB redundant slot count must match" - def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): - return [ - [ - (layer_index, plan, buffer_index, self.staging[buffer_index]) - for buffer_index, (layer_index, plan) in enumerate( - layer_plans[batch_start : batch_start + self.staging_depth] - ) - ] - for batch_start in range(0, len(layer_plans), self.staging_depth) - ] - - def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepared_batches) -> None: - if self._thread is not None and self._thread.is_alive(): - raise RuntimeError("EPLB transfer is already in flight") - expected_batch_count = (len(layer_plans) + self.staging_depth - 1) // self.staging_depth - if len(prepared_batches) != expected_batch_count: - raise ValueError("EPLB prepared batch count does not match layer-plan batches") - # 预构造批次与描述符一起传入,避免推理线程重建。 - self._start_transfer_generation() + @staticmethod + def _group_steps_by_source(plan: Sequence[TransferStep]): + grouped = defaultdict(list) + for step in plan: + grouped[(step.src_rank, step.src_local_row)].append(step) + return [(source, grouped[source]) for source in sorted(grouped)] + + def _copy_layer(self, layer_index: int, plan: Sequence[TransferStep]) -> None: + live_tensors = self.live[layer_index] + for (src_rank, src_local_row), steps in self._group_steps_by_source(plan): + if self.global_rank == src_rank: + with torch.cuda.stream(self._copy_stream): + for (_, live), (_, pinned) in zip(live_tensors, self.pinned_rows): + pinned.copy_(live[src_local_row], non_blocking=True) + # Gloo must not read the CPU row before the device-to-host copy completes. + self._copy_stream.synchronize() + for _, pinned in self.pinned_rows: + dist.broadcast(pinned, src=src_rank, group=self.transfer_group) + + dst_slots = sorted({step.dst_slot for step in steps if step.dst_rank == self.global_rank}) + if dst_slots: + with torch.cuda.stream(self._copy_stream): + for (_, staging), (_, pinned) in zip(self.staging, self.pinned_rows): + for dst_slot in dst_slots: + staging[dst_slot].copy_(pinned, non_blocking=True) + self._copy_stream.synchronize() + + def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]) -> None: + if self._thread is not None: + raise RuntimeError("EPLB transfer has not been finished") self._error = None with self._pending_lock: self._pending.clear() @@ -178,33 +186,21 @@ def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepa def worker() -> None: try: torch.cuda.set_device(self.device) - if not layer_plans: - self._finish_transfer_generation() - for batch_index, (batch, prepared_batch) in enumerate(prepared_batches): - batch_start = batch_index * self.staging_depth - for layer_index, plan, buffer_index, _ in batch: - release = self._release[buffer_index] - # A buffer cannot be reused until its prior committed rows are no longer read by CUDA. - release.wait() - release.clear() - if self._consumed_recorded[buffer_index]: - self._consumed_events[buffer_index].synchronize() - self._changed_dst_slots[buffer_index] = tuple( - step.dst_slot for step in plan if step.dst_rank == self.global_rank - ) - if batch_start > 0 and self._needs_staging_reuse_barrier: - # All destinations must finish consuming the prior IPC staging generation - # before a source can reuse the peer buffer for this batch. - dist.barrier(group=self.transfer_group) - self._copy_batch(batch, prepared_batch) - if batch_start + self.staging_depth >= len(layer_plans): - self._finish_transfer_generation() + for layer_index, plan in layer_plans: + self._release.wait() + self._release.clear() + if self._consumed_recorded: + self._consumed_event.synchronize() + self._changed_dst_slots = tuple( + sorted({step.dst_slot for step in plan if step.dst_rank == self.global_rank}) + ) + self._copy_layer(layer_index, plan) with self._pending_lock: - self._pending.extend((layer_index, buffer_index) for layer_index, _, buffer_index, _ in batch) + self._pending.append((layer_index, 0)) except BaseException as exc: self._error = exc - self._thread = threading.Thread(target=worker, name=f"eplb-{self.backend}", daemon=True) + self._thread = threading.Thread(target=worker, name="eplb-pin-memory", daemon=True) self._thread.start() def pending_layers(self): @@ -214,26 +210,22 @@ def pending_layers(self): return list(self._pending) def commit(self, layer_index: int, buffer_index: int, post_copy=None) -> None: + if buffer_index != 0: + raise RuntimeError("Pinned-memory EPLB uses a single staging buffer") with self._pending_lock: if not self._pending or self._pending[0] != (layer_index, buffer_index): raise RuntimeError("EPLB commit does not match the pending FIFO") self._pending.popleft() - changed_dst_slots = self._changed_dst_slots[buffer_index] - for (_, live), (_, staging) in zip(self.live[layer_index], self.staging[buffer_index]): - _commit_staging_rows( - live, - staging, - self.num_experts_per_rank, - changed_dst_slots, - ) + changed_dst_slots = self._changed_dst_slots + for (_, live), (_, staging) in zip(self.live[layer_index], self.staging): + _commit_staging_rows(live, staging, self.num_experts_per_rank, changed_dst_slots) if post_copy is not None: post_copy() - self._consumed_events[buffer_index].record(torch.cuda.current_stream()) - self._consumed_recorded[buffer_index] = True - self._release[buffer_index].set() + self._consumed_event.record(torch.cuda.current_stream()) + self._consumed_recorded = True + self._release.set() def finish(self) -> None: - """Wait for the released migration worker to exit before another rebalance.""" thread = self._thread if thread is None: return @@ -243,375 +235,6 @@ def finish(self) -> None: raise RuntimeError("EPLB migration worker failed") from self._error -class NixlEPLBTransfer(_EPLBTransferBase): - """GPU-direct UCX/NIXL EPLB transfer. Initialization errors are fatal.""" - - backend = "nixl" - _DEFAULT_UCX_TLS = "self,sm,cuda_ipc,cuda_copy,rc_x" - - @dataclass - class _PreparedBatch: - remote_entries: Dict[int, list] - push_batch: "_PreparedCudaMemcpyBatch | None" - - def __init__(self, weights, transfer_group, global_rank, world_size): - # Reuse at most eight layer buffers to bound EPLB staging memory. - self.staging_depth = min(EPLB_MAX_STAGING_DEPTH, len(weights)) - super().__init__(weights, transfer_group, global_rank, world_size) - self._nixl_agent = None - self._registered_descs = None - self._remote_agents: Dict[int, str] = {} - self._remote_layouts = {} - self._xfer_cache = {} - self._used_xfer_cache_keys = set() - self._ipc_staging = {} - self._same_node_ranks = set() - self._cross_node_ranks = set() - self._push_stream = torch.cuda.Stream(device=self.device) - self._batch_memcpy = _CudaBatchMemcpy() - try: - self._init_ipc_metadata() - self._init_push_layouts() - if self._cross_node_ranks: - os.environ.setdefault("UCX_TLS", self._DEFAULT_UCX_TLS) - try: - import nixl - except Exception as exc: - raise RuntimeError("NIXL EPLB backend requires the nixl package for cross-node transfer") from exc - agent_name = f"lightllm-eplb-{socket.gethostname()}-{os.getpid()}-rank-{global_rank}" - config = nixl.nixl_agent_config(enable_prog_thread=True, enable_listen_thread=False, backends=["UCX"]) - self._nixl_agent = nixl.nixl_agent(agent_name, config) - reg_tensors = [tensor for layer in self.live for _, tensor in layer] + [ - tensor for staging in self.staging for _, tensor in staging - ] - self._registered_descs = self._nixl_agent.get_reg_descs(reg_tensors) - self._nixl_agent.register_memory(self._registered_descs, backends=["UCX"]) - self._init_remote_metadata() - except Exception as exc: - self.shutdown() - if isinstance(exc, RuntimeError): - raise - raise RuntimeError("NIXL EPLB initialization failed") from exc - - def _local_layout(self): - return [ - [(name, tensor.data_ptr(), tensor.get_device(), tensor[0].nbytes) for name, tensor in layer] - for layer in self.live - ] - - def _init_ipc_metadata(self) -> None: - hostnames = [None] * self.world_size - dist.all_gather_object(hostnames, socket.gethostname(), group=self.transfer_group) - local_hostname = hostnames[self.global_rank] - self._needs_staging_reuse_barrier = len(set(hostnames)) < len(hostnames) - self._same_node_ranks = {rank for rank, hostname in enumerate(hostnames) if hostname == local_hostname} - self._cross_node_ranks = set(range(self.world_size)) - self._same_node_ranks - from lightllm.server.router.model_infer.mode_backend.pd.p2p_fix import ( - p2p_fix_rebuild_cuda_tensor, - reduce_tensor, - ) - - exports = {} - for target_rank in self._same_node_ranks - {self.global_rank}: - exports[target_rank] = { - "staging": [ - [(name, tuple(tensor.shape), tensor.dtype, reduce_tensor(tensor)[1]) for name, tensor in staging] - for staging in self.staging - ], - } - all_exports = [None] * self.world_size - dist.all_gather_object(all_exports, exports, group=self.transfer_group) - - torch.cuda.set_device(self.device) - for dst_rank in self._same_node_ranks - {self.global_rank}: - metadata = all_exports[dst_rank].get(self.global_rank) - if metadata is None or len(metadata["staging"]) != self.staging_depth: - raise RuntimeError(f"NIXL IPC destination rank {dst_rank} has incompatible staging metadata") - rebuilt_staging = [] - for remote_staging, local_staging in zip(metadata["staging"], self.staging): - if len(remote_staging) != len(local_staging): - raise RuntimeError(f"NIXL IPC destination rank {dst_rank} staging tensor count mismatch") - rebuilt = [] - for (name, shape, dtype, args), (local_name, local_tensor) in zip(remote_staging, local_staging): - if name != local_name or shape != tuple(local_tensor.shape) or dtype != local_tensor.dtype: - raise RuntimeError(f"NIXL IPC destination rank {dst_rank} staging layout mismatch for {name}") - tensor = p2p_fix_rebuild_cuda_tensor(*args) - if tuple(tensor.shape) != shape or tensor.dtype != dtype or tensor.device != local_tensor.device: - raise RuntimeError( - f"NIXL IPC destination rank {dst_rank} staging rebuild validation failed for {name}" - ) - rebuilt.append((name, tensor)) - rebuilt_staging.append(rebuilt) - self._ipc_staging[dst_rank] = rebuilt_staging - - def _init_remote_metadata(self) -> None: - metadata = self._nixl_agent.get_agent_metadata() - all_metadata = [None] * self.world_size - all_layouts = [None] * self.world_size - dist.all_gather_object(all_metadata, metadata, group=self.transfer_group) - dist.all_gather_object(all_layouts, self._local_layout(), group=self.transfer_group) - for rank in self._cross_node_ranks: - layout = all_layouts[rank] - if len(layout) != len(self.live): - raise RuntimeError(f"NIXL remote rank {rank} has incompatible layer layout") - self._remote_agents[rank] = self._nixl_agent.add_remote_agent(all_metadata[rank]) - self._remote_layouts[rank] = layout - - def _wait_xfers(self, xfers) -> None: - pending = [] - for item in xfers: - state = self._nixl_agent.transfer(item[2]) - if state == "ERR": - raise RuntimeError("NIXL READ post failed") - if state == "PROC": - pending.append(item) - while pending: - remaining = [] - for item in pending: - state = self._nixl_agent.check_xfer_state(item[2]) - if state == "ERR": - raise RuntimeError("NIXL READ transfer failed") - if state != "DONE": - remaining.append(item) - pending = remaining - - def _release_xfers(self, xfers) -> None: - unreleased = [] - errors = [] - for local_dlist, remote_dlist, xfer in xfers: - remaining = [local_dlist, remote_dlist, xfer] - for remaining_index, handle, release in ( - (2, xfer, self._nixl_agent.release_xfer_handle), - (1, remote_dlist, self._nixl_agent.release_dlist_handle), - (0, local_dlist, self._nixl_agent.release_dlist_handle), - ): - if handle is not None: - try: - release(handle) - except Exception as exc: - errors.append(exc) - else: - remaining[remaining_index] = None - if any(handle is not None for handle in remaining): - unreleased.append(tuple(remaining)) - if errors: - error = RuntimeError("NIXL transfer handle release failed") - error.unreleased_xfers = unreleased - raise error from errors[0] - - @staticmethod - def _contiguous_runs(steps): - ordered = sorted(steps, key=lambda step: (step.src_local_row, step.dst_slot)) - runs = [] - for step in ordered: - if ( - runs - and step.src_local_row == runs[-1][-1].src_local_row + 1 - and step.dst_slot == runs[-1][-1].dst_slot + 1 - ): - runs[-1].append(step) - else: - runs.append([step]) - return runs - - @staticmethod - def _remote_read_cache_key(src_rank: int, entries): - return ( - src_rank, - tuple( - ( - layer_index, - tuple((step.src_local_row, step.dst_slot) for step in run), - tuple(tensor.data_ptr() for _, tensor in staging), - ) - for layer_index, run, staging in entries - ), - ) - - def _init_push_layouts(self) -> None: - self._live_row_layout = [ - [(name, tensor.data_ptr(), tensor[0].nbytes) for name, tensor in layer] for layer in self.live - ] - reference = [(name, row_nbytes) for name, _, row_nbytes in self._live_row_layout[0]] - self._push_staging_row_layout = {} - for dst_rank in self._same_node_ranks: - layouts = [] - for buffer_index in range(self.staging_depth): - staging = ( - self.staging[buffer_index] - if dst_rank == self.global_rank - else self._ipc_staging[dst_rank][buffer_index] - ) - layout = [(name, tensor.data_ptr(), tensor[0].nbytes) for name, tensor in staging] - if [(name, row_nbytes) for name, _, row_nbytes in layout] != reference: - raise RuntimeError("NIXL source-push staging row layout mismatch") - layouts.append(layout) - self._push_staging_row_layout[dst_rank] = layouts - - def _prepare_batch(self, batch): - remote_entries = defaultdict(list) - push_descriptors = [] - for layer_index, plan, buffer_index, staging in batch: - steps_by_source = defaultdict(list) - by_destination = defaultdict(list) - for step in plan: - if step.dst_rank == self.global_rank and step.src_rank not in self._same_node_ranks: - steps_by_source[step.src_rank].append(step) - if step.src_rank == self.global_rank and step.dst_rank in self._same_node_ranks: - by_destination[step.dst_rank].append(step) - for src_rank, steps in steps_by_source.items(): - remote_entries[src_rank].extend((layer_index, run, staging) for run in self._contiguous_runs(steps)) - source_layout = self._live_row_layout[layer_index] - for dst_rank, steps in by_destination.items(): - destination_layout = self._push_staging_row_layout[dst_rank][buffer_index] - for run in self._contiguous_runs(steps): - first = run[0] - run_len = len(run) - for (_, source_ptr, row_nbytes), (_, destination_ptr, _) in zip(source_layout, destination_layout): - push_descriptors.append( - ( - source_ptr + first.src_local_row * row_nbytes, - destination_ptr + first.dst_slot * row_nbytes, - run_len * row_nbytes, - ) - ) - push_batch = self._batch_memcpy.prepare(push_descriptors) if push_descriptors else None - return self._PreparedBatch(dict(remote_entries), push_batch) - - def prepare_transfer(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): - return [(batch, self._prepare_batch(batch)) for batch in self._make_batches(layer_plans)] - - def _get_remote_read(self, src_rank: int, entries): - cache_key = self._remote_read_cache_key(src_rank, entries) - cached = self._xfer_cache.get(cache_key) - if cached is not None: - self._used_xfer_cache_keys.add(cache_key) - return cached - local_descs = [] - remote_descs = [] - local_dlist = remote_dlist = xfer = None - try: - for layer_index, run, staging in entries: - remote_layer = self._remote_layouts[src_rank][layer_index] - if len(remote_layer) != len(staging): - raise RuntimeError(f"NIXL remote rank {src_rank} has incompatible layer layout") - first = run[0] - run_len = len(run) - for tensor_index, (_, staging_tensor) in enumerate(staging): - name, remote_ptr, remote_device, remote_nbytes = remote_layer[tensor_index] - if ( - name != self.live[layer_index][tensor_index][0] - or remote_nbytes != staging_tensor[first.dst_slot].nbytes - ): - raise RuntimeError(f"NIXL remote rank {src_rank} descriptor range mismatch") - local_descs.append( - ( - staging_tensor[first.dst_slot].data_ptr(), - run_len * remote_nbytes, - staging_tensor.get_device(), - ) - ) - remote_descs.append( - (remote_ptr + first.src_local_row * remote_nbytes, run_len * remote_nbytes, remote_device) - ) - local_dlist = self._nixl_agent.prep_xfer_dlist( - "NIXL_INIT_AGENT", self._nixl_agent.get_xfer_descs(local_descs, "VRAM"), backends=["UCX"] - ) - remote_dlist = self._nixl_agent.prep_xfer_dlist( - self._remote_agents[src_rank], self._nixl_agent.get_xfer_descs(remote_descs, "VRAM"), backends=["UCX"] - ) - xfer = self._nixl_agent.make_prepped_xfer( - "READ", - local_dlist, - list(range(len(local_descs))), - remote_dlist, - list(range(len(remote_descs))), - backends=["UCX"], - ) - selected_backend = self._nixl_agent.query_xfer_backend(xfer) - if selected_backend != "UCX": - raise RuntimeError("NIXL EPLB READ did not select UCX") - self._xfer_cache[cache_key] = (local_dlist, remote_dlist, xfer) - self._used_xfer_cache_keys.add(cache_key) - return self._xfer_cache[cache_key] - except Exception: - self._release_xfers([(local_dlist, remote_dlist, xfer)]) - raise - - def _copy_batch(self, batch, prepared_batch) -> None: - if prepared_batch.push_batch is not None: - self._batch_memcpy.enqueue(prepared_batch.push_batch, self._push_stream.cuda_stream) - xfers = [ - self._get_remote_read(src_rank, entries) for src_rank, entries in prepared_batch.remote_entries.items() - ] - self._wait_xfers(xfers) - self._push_stream.synchronize() - # Before a rank publishes this batch it has completed its outgoing source-pushes and - # incoming UCX READs. The manager's global MIN-ready gate therefore means all transfers - # are complete before any rank commits, without a destination-side GPU wait. - - def _start_transfer_generation(self) -> None: - self._used_xfer_cache_keys.clear() - - def _finish_transfer_generation(self) -> None: - errors = [] - for cache_key in set(self._xfer_cache) - self._used_xfer_cache_keys: - xfer = self._xfer_cache[cache_key] - try: - self._release_xfers([xfer]) - except Exception as exc: - unreleased = getattr(exc, "unreleased_xfers", None) - if unreleased: - self._xfer_cache[cache_key] = unreleased[0] - errors.append(exc) - else: - del self._xfer_cache[cache_key] - if errors: - raise RuntimeError("NIXL EPLB cache eviction failed") from errors[0] - - def shutdown(self) -> None: - agent = self._nixl_agent - errors = [] - getattr(self, "_used_xfer_cache_keys", set()).clear() - if agent is not None: - for cache_key, xfer in list(self._xfer_cache.items()): - try: - self._release_xfers([xfer]) - except Exception as exc: - unreleased = getattr(exc, "unreleased_xfers", None) - if unreleased: - self._xfer_cache[cache_key] = unreleased[0] - errors.append(exc) - else: - del self._xfer_cache[cache_key] - if errors: - raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] - for remote_name in list(self._remote_agents.values()): - if agent is not None: - try: - agent.remove_remote_agent(remote_name) - except Exception as exc: - errors.append(exc) - self._remote_agents.clear() - self._remote_layouts.clear() - if agent is not None and self._registered_descs is not None: - try: - agent.deregister_memory(self._registered_descs, backends=["UCX"]) - except Exception as exc: - errors.append(exc) - self._registered_descs = None - self._nixl_agent = None - getattr(self, "_ipc_staging", {}).clear() - if errors: - raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] - - def __del__(self): - try: - self.shutdown() - except Exception: - pass - - def _commit_staging_rows( live: torch.Tensor, staging: torch.Tensor, @@ -632,135 +255,3 @@ def _commit_staging_rows( ) if dst_slot is not None: run_start = previous = dst_slot - - -class _CudaMemLocation(ctypes.Structure): - _fields_ = [("type", ctypes.c_int), ("id", ctypes.c_int)] - - -class _CudaMemcpyAttributes(ctypes.Structure): - _fields_ = [ - ("srcAccessOrder", ctypes.c_int), - ("srcLocHint", _CudaMemLocation), - ("dstLocHint", _CudaMemLocation), - ("flags", ctypes.c_uint), - ] - - -@dataclass -class _PreparedCudaMemcpyBatch: - """Host-side arrays retained for one cudaMemcpyBatchAsync submission.""" - - dsts: object - srcs: object - sizes: object - attrs: _CudaMemcpyAttributes - attrs_idxs: object - count: int - - -class _CudaBatchMemcpy: - """CUDA 13.x ``cudaMemcpyBatchAsync`` binding for EPLB source-push.""" - - _SRC_ACCESS_ORDER_STREAM = 1 - _PREFER_OVERLAP_WITH_COMPUTE = 1 - _CUDA_13_0 = 13000 - _CUDA_14_0 = 14000 - - def __init__(self, library=None): - if library is None: - path = self._find_loaded_cudart() - if path is None: - raise RuntimeError( - "NIXL same-node source-push requires CUDA Runtime 13.x cudaMemcpyBatchAsync; " - "libcudart.so.13 is not loaded" - ) - try: - library = ctypes.CDLL(path) - except OSError as exc: - raise RuntimeError(f"cannot load libcudart: {exc}") from exc - - try: - runtime_get_version = library.cudaRuntimeGetVersion - self._batch_async = library.cudaMemcpyBatchAsync - self._get_error_string = library.cudaGetErrorString - except AttributeError as exc: - raise RuntimeError("cudaMemcpyBatchAsync is unavailable") from exc - - runtime_get_version.restype = ctypes.c_int - runtime_get_version.argtypes = [ctypes.POINTER(ctypes.c_int)] - self._get_error_string.restype = ctypes.c_char_p - self._get_error_string.argtypes = [ctypes.c_int] - runtime_version = ctypes.c_int() - result = runtime_get_version(ctypes.byref(runtime_version)) - if result != 0: - raise RuntimeError(f"cudaRuntimeGetVersion failed with CUDA error {result}") - if not self._CUDA_13_0 <= runtime_version.value < self._CUDA_14_0: - raise RuntimeError( - f"cudaMemcpyBatchAsync requires CUDA Runtime 13.x (13.0 ABI), found {runtime_version.value}" - ) - - pointer_array = ctypes.POINTER(ctypes.c_void_p) - self._batch_async.restype = ctypes.c_int - self._batch_async.argtypes = [ - pointer_array, - pointer_array, - ctypes.POINTER(ctypes.c_size_t), - ctypes.c_size_t, - ctypes.POINTER(_CudaMemcpyAttributes), - ctypes.POINTER(ctypes.c_size_t), - ctypes.c_size_t, - ctypes.c_void_p, - ] - - @staticmethod - def prepare(copies: Iterable[Tuple[int, int, int]]) -> _PreparedCudaMemcpyBatch: - copies = tuple(copies) - if not copies: - raise ValueError("cudaMemcpyBatchAsync requires at least one copy") - for src, dst, size in copies: - if not src or not dst or size <= 0: - raise ValueError("cudaMemcpyBatchAsync requires non-null pointers and positive sizes") - count = len(copies) - dsts = (ctypes.c_void_p * count)(*(dst for _, dst, _ in copies)) - srcs = (ctypes.c_void_p * count)(*(src for src, _, _ in copies)) - sizes = (ctypes.c_size_t * count)(*(size for _, _, size in copies)) - attrs = _CudaMemcpyAttributes() - attrs.srcAccessOrder = _CudaBatchMemcpy._SRC_ACCESS_ORDER_STREAM - attrs.flags = _CudaBatchMemcpy._PREFER_OVERLAP_WITH_COMPUTE - attrs_idxs = (ctypes.c_size_t * 1)(0) - return _PreparedCudaMemcpyBatch(dsts, srcs, sizes, attrs, attrs_idxs, count) - - def enqueue(self, prepared: _PreparedCudaMemcpyBatch, stream: int) -> None: - result = self._batch_async( - prepared.dsts, - prepared.srcs, - prepared.sizes, - prepared.count, - ctypes.byref(prepared.attrs), - prepared.attrs_idxs, - 1, - ctypes.c_void_p(stream), - ) - if result != 0: - message = self._get_error_string(result) - error = message.decode("utf-8") if message else f"CUDA error {result}" - raise RuntimeError(f"cudaMemcpyBatchAsync failed: {error}") - - @staticmethod - def _find_loaded_cudart() -> Optional[str]: - """Return a mapped CUDA 13 runtime without loading CUDA as a side effect.""" - try: - with open("/proc/self/maps") as maps: - for line in maps: - if "libcudart" not in line: - continue - path_start = line.find("/") - if path_start < 0: - continue - path = line[path_start:].strip().removesuffix(" (deleted)") - if re.search(r"libcudart[^/]*\.so\.13(?:\D|$)", path): - return path - except OSError: - pass - return None diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index da4324e382..131c853f89 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1,9 +1,6 @@ -import builtins -import io import threading import time -from collections import deque -from contextlib import contextmanager +from contextlib import nullcontext from types import SimpleNamespace import pytest @@ -42,8 +39,8 @@ ) from lightllm.common.eplb_utils import extract_eplb_expert_tensors from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + PinnedMemoryEPLBTransfer, TransferStep, - _CudaBatchMemcpy, _commit_staging_rows, align_target_placement, build_transfer_plan, @@ -288,8 +285,8 @@ def test_eplb_redundant_experts_default_to_disabled(): @pytest.mark.parametrize( ("num_logical_experts", "num_ranks", "num_redundant_experts_per_rank", "expected"), [ - (8, 4, 2, [[0, 1, 1, 1], [2, 3, 3, 3], [4, 5, 5, 5], [6, 7, 7, 7]]), - (6, 3, 4, [[0, 1, 1, 1, 1, 1], [2, 3, 3, 3, 3, 3], [4, 5, 5, 5, 5, 5]]), + (8, 4, 2, [[0, 1, 2, 3], [2, 3, 4, 5], [4, 5, 6, 7], [6, 7, 0, 1]]), + (6, 3, 4, [[0, 1, 2, 3, 4, 5], [2, 3, 4, 5, 0, 1], [4, 5, 0, 1, 2, 3]]), ], ) def test_build_initial_local_expert_ids( @@ -307,6 +304,41 @@ def test_build_initial_local_expert_ids( assert actual == expected +def test_build_initial_local_expert_ids_rejects_local_or_duplicate_replicas(): + with pytest.raises(AssertionError): + build_initial_local_expert_ids(8, 4, 7) + + +def test_fused_moe_loads_default_replicas_into_their_physical_rows(): + weight = object.__new__(fused_weight_module.FusedMoeWeight) + weight.lock = threading.Lock() + loaded = [] + + def load_weight(expert, local, _weights): + loaded.append(("weight", expert, local)) + + def load_scale(expert, local, _weights): + loaded.append(("scale", expert, local)) + + def load_zero_point(expert, local, _weights): + loaded.append(("zero", expert, local)) + + weight._load_expert = load_weight + weight._load_expert_scale = load_scale + weight._load_expert_zero_point = load_zero_point + local_logic_expert_ids_list = build_initial_local_expert_ids(8, 4, 2)[0] + + weight._load_weight(local_logic_expert_ids_list, {}) + + assert local_logic_expert_ids_list == [0, 1, 2, 3] + assert [entry for entry in loaded if entry[0] == "weight"] == [ + ("weight", 0, 0), + ("weight", 1, 1), + ("weight", 2, 2), + ("weight", 3, 3), + ] + + def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): expert_load = ( torch.tensor( @@ -899,7 +931,7 @@ def test_align_target_placement_keeps_retained_experts_in_live_slots(): assert torch.equal(canonical, torch.tensor([[6, 5], [0, 7], [0, 1], [2, 3]])) -def test_align_target_placement_replaces_duplicate_initial_placeholders(): +def test_align_target_placement_replaces_duplicate_current_replicas(): current = torch.tensor([[1, 1], [3, 3]]) target = torch.tensor([[1, 2], [3, 0]]) @@ -1539,14 +1571,6 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c manager._evaluation_result = None manager._evaluation_error = None manager._continuous_collection_end_step = None - prepare_calls = [] - prepared_batches = object() - - def prepare_transfer(layer_plans): - prepare_calls.append(layer_plans) - return prepared_batches - - manager.transfer = SimpleNamespace(prepare_transfer=prepare_transfer) manager._collect_local_samples = lambda: torch.full((1, 3, 4), 100, dtype=torch.int64) planned_placement = torch.tensor( [ @@ -1584,8 +1608,6 @@ def build_maps_for_layers(*args, **kwargs): metadata = manager._evaluation_result["metadata"] assert metadata[1] is None assert [layer_index for layer_index, _plan in manager._evaluation_result["layer_plans"]] == [0, 2] - assert prepare_calls == [manager._evaluation_result["layer_plans"]] - assert manager._evaluation_result["prepared_batches"] is prepared_batches for layer_index in (0, 2): item = metadata[layer_index] expected = build_logical_to_physical_map( @@ -1596,57 +1618,6 @@ def build_maps_for_layers(*args, **kwargs): assert torch.equal(item, torch.tensor(expected, dtype=torch.int32)) -def test_manager_preparation_error_is_saved_as_evaluation_error(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "fuse_moe_impl": _test_moe_impl( - eplb=True, - route_counter=torch.zeros((2,), dtype=torch.int64), - num_logical_experts=2, - world_size=1, - ) - }, - )() - ] - manager._eplb_impls = [manager.weights[0].fuse_moe_impl] - manager.global_rank = 0 - manager.world_size = 1 - manager.node_world_size = 1 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.num_logical_experts = 2 - manager.current_placement = torch.tensor([[[0]]], dtype=torch.int64) - manager.evaluation_group = object() - manager._evaluation_lock = threading.Lock() - manager._evaluation_result = None - manager._evaluation_error = None - manager._continuous_collection_end_step = None - manager._collect_local_samples = lambda: torch.ones((1, 1, 2), dtype=torch.int64) - manager._plan_and_broadcast = lambda _global_load: { - "kind": "planned", - "placement": torch.tensor([[[0]]], dtype=torch.int64), - "improved": torch.tensor([True]), - } - - def fail_prepare(_layer_plans): - raise RuntimeError("prepare failed") - - manager.transfer = SimpleNamespace(prepare_transfer=fail_prepare) - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) - monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) - monkeypatch.setattr(manager_module, "build_transfer_plan", lambda *_args, **_kwargs: []) - - manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) - - assert manager._evaluation_result is None - assert isinstance(manager._evaluation_error, RuntimeError) - assert str(manager._evaluation_error) == "prepare failed" - - def test_decode_dispatch_uses_physical_ids_and_total_expert_count(monkeypatch): class Buffer: def low_latency_dispatch(self, **kwargs): @@ -1891,7 +1862,7 @@ def cpu_zeros(*shape, **kwargs): assert impl.num_total_physical_experts == 6 assert impl.route_counter.shape == (4,) assert impl.recording - assert impl.local_logics_expert_ids_list == [0, 1, 1] + assert impl.local_logics_expert_ids_list == [0, 1, 2] assert not hasattr(impl, "initial_local_expert_ids_by_rank") assert not hasattr(impl, "expert_parallel_state") @@ -2340,139 +2311,6 @@ def finish(self): g_infer_context.overlap_stream = original_overlap_stream -def test_transfer_ring_reuses_a_buffer_only_after_commit_and_consumption(monkeypatch): - operations = [] - - class Event: - def __init__(self): - self.recorded = 0 - self.synchronized = 0 - - def record(self, stream): - self.recorded += 1 - - def synchronize(self): - self.synchronized += 1 - operations.append("consumed synchronize") - - transfer = object.__new__(transfer_module._EPLBTransferBase) - transfer.backend = "test" - transfer.device = torch.device("cuda", 0) - transfer.staging_depth = 2 - transfer.staging = [[], []] - transfer.live = [[], [], []] - transfer.num_experts_per_rank = 0 - transfer._release = [threading.Event(), threading.Event()] - for release in transfer._release: - release.set() - transfer._consumed_events = [Event(), Event()] - transfer._consumed_recorded = [False, False] - transfer._changed_dst_slots = [(), ()] - transfer._pending = deque() - transfer._pending_lock = threading.Lock() - transfer._error = None - transfer._thread = None - transfer._needs_staging_reuse_barrier = True - transfer._start_transfer_generation = lambda: None - transfer._finish_transfer_generation = lambda: None - transfer.transfer_group = "transfer-group" - copied = [] - - def copy_batch(batch, _prepared_batch): - for layer, _plan, _buffer, _staging in batch: - copied.append(layer) - operations.append(("copy", layer)) - - transfer._copy_batch = copy_batch - monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda device: None) - monkeypatch.setattr(transfer_module.torch.cuda, "Event", Event) - monkeypatch.setattr(transfer_module.torch.cuda, "current_stream", lambda: object()) - monkeypatch.setattr( - transfer_module.dist, - "barrier", - lambda **kwargs: operations.append(("barrier", kwargs["group"])), - ) - - plans = [(0, []), (1, []), (2, [])] - prepared_batches = [(batch, None) for batch in transfer._make_batches(plans)] - monkeypatch.setattr( - transfer, - "_make_batches", - lambda _plans: pytest.fail("start must reuse prepared batches"), - ) - transfer.start(plans, prepared_batches) - deadline = time.monotonic() + 2 - while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: - time.sleep(0.001) - assert transfer.pending_layers() == [(0, 0), (1, 1)] - assert copied == [0, 1] - assert operations == [("copy", 0), ("copy", 1)] - - transfer.commit(0, 0) - while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: - time.sleep(0.001) - assert transfer.pending_layers() == [(1, 1), (2, 0)] - assert copied == [0, 1, 2] - assert transfer._consumed_events[0].synchronized == 1 - assert operations == [ - ("copy", 0), - ("copy", 1), - "consumed synchronize", - ("barrier", "transfer-group"), - ("copy", 2), - ] - transfer.commit(1, 1) - transfer.commit(2, 0) - transfer.finish() - - -def test_transfer_finalization_failure_stays_in_worker_and_success_finalizes_once( - monkeypatch, -): - def make_transfer(finalize): - transfer = object.__new__(transfer_module._EPLBTransferBase) - transfer.backend = "test" - transfer.device = torch.device("cuda", 0) - transfer.global_rank = 0 - transfer.staging_depth = 1 - transfer.staging = [[]] - transfer._release = [threading.Event()] - transfer._release[0].set() - transfer._consumed_events = [object()] - transfer._consumed_recorded = [False] - transfer._changed_dst_slots = [()] - transfer._pending = deque() - transfer._pending_lock = threading.Lock() - transfer._error = None - transfer._thread = None - transfer._needs_staging_reuse_barrier = False - transfer._copy_batch = lambda _batch, _prepared_batch: None - transfer._start_transfer_generation = lambda: None - transfer._finish_transfer_generation = finalize - return transfer - - monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) - - finalized_before_publish = [] - success = make_transfer(lambda: finalized_before_publish.append(len(success._pending))) - success.start([(0, [])], [([(0, [], 0, [])], None)]) - success.finish() - assert finalized_before_publish == [0] - assert success.pending_layers() == [(0, 0)] - - failed = make_transfer(lambda: (_ for _ in ()).throw(RuntimeError("cache boom"))) - failed.start([(0, [])], [([(0, [], 0, [])], None)]) - failed._thread.join() - assert list(failed._pending) == [] - with pytest.raises(RuntimeError, match="EPLB migration worker failed") as exc_info: - failed.pending_layers() - assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "cache boom" - with pytest.raises(RuntimeError, match="EPLB migration worker failed"): - failed.finish() - assert failed._thread is None - - def test_manager_rearms_after_rebalance_for_interval_one(): recording_calls = [] manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) @@ -2773,18 +2611,16 @@ def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): manager.transfer = type( "Transfer", (), - {"start": lambda self, plans, prepared_batches: setattr(self, "started", (plans, prepared_batches))}, + {"start": lambda self, plans: setattr(self, "started", plans)}, )() manager._reset_route_counters = lambda: None - prepared_batches = [object()] manager._start_rebalance( { "placement": torch.zeros((1, 1, 1), dtype=torch.int64), "improved": torch.tensor([True]), "metadata": [None], "layer_plans": [(0, object())], - "prepared_batches": prepared_batches, "before": {"max": 1.0, "p95": 1.0}, "after": {"max": 1.0, "p95": 1.0}, "model_imbalance_ratio": 1.0, @@ -2797,18 +2633,7 @@ def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): assert manager.sampling_interval == 20 assert manager.in_flight assert manager._continuous_collection_start_step is None - assert len(manager.transfer.started[0]) == 1 - assert manager.transfer.started[1] is prepared_batches - - -def test_transfer_start_rejects_prepared_batches_with_wrong_batch_count(): - transfer = object.__new__(transfer_module._EPLBTransferBase) - transfer._thread = None - transfer.staging_depth = 2 - transfer.staging = [object(), object()] - - with pytest.raises(ValueError, match="prepared batch count"): - transfer.start([(0, []), (1, []), (2, [])], prepared_batches=[object()]) + assert len(manager.transfer.started) == 1 def test_first_rebalance_completion_switches_to_four_step_sparse_window(monkeypatch): @@ -2882,596 +2707,108 @@ def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): assert calls == ["evaluation"] -def test_nixl_descriptor_runs_merge_only_jointly_contiguous_source_and_destination_rows(): - steps = [ +def test_pinned_transfer_groups_sources_in_collective_order(): + plan = [ TransferStep(0, 2, 1, 3), - TransferStep(0, 0, 1, 1), - TransferStep(0, 1, 1, 2), - TransferStep(0, 4, 1, 7), - ] - runs = transfer_module.NixlEPLBTransfer._contiguous_runs(steps) - assert [[(step.dst_slot, step.src_local_row) for step in run] for run in runs] == [ - [(0, 1), (1, 2), (2, 3)], - [(4, 7)], - ] - - -def test_nixl_prepare_batch_compiles_hot_path_without_tensor_views(monkeypatch): - class BatchMemcpy: - def __init__(self): - self.prepared = [] - self.enqueued = [] - - def prepare(self, descriptors): - descriptor = tuple(descriptors) - self.prepared.append(descriptor) - return descriptor - - def enqueue(self, descriptor, stream): - self.enqueued.append((descriptor, stream)) - - stream = SimpleNamespace(cuda_stream=123, synchronize=lambda: None) - batch_memcpy = BatchMemcpy() - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._push_stream = stream - transfer._batch_memcpy = batch_memcpy - transfer.global_rank = 0 - transfer._same_node_ranks = {0, 1, 2} - transfer._live_row_layout = [ - [ - ("w13.weight", 1000, 32), - ("w13.weight_scale", 2000, 32), - ("w2.weight", 3000, 32), - ] - ] - transfer._push_staging_row_layout = { - 1: [ - [ - ("w13.weight", 4000, 32), - ("w13.weight_scale", 5000, 32), - ("w2.weight", 6000, 32), - ] - ], - 2: [ - [ - ("w13.weight", 7000, 32), - ("w13.weight_scale", 8000, 32), - ("w2.weight", 9000, 32), - ] - ], - } - transfer._get_remote_read = lambda *_args: None - transfer._wait_xfers = lambda _xfers: None - - run_a = [TransferStep(1, 3, 0, 5), TransferStep(1, 4, 0, 6)] - run_b = [TransferStep(2, 1, 0, 2)] - local_inbound = TransferStep(0, 0, 1, 0) - remote_inbound = TransferStep(0, 1, 3, 2) - staging = object() - batch = [(0, run_a + run_b + [local_inbound, remote_inbound], 0, staging)] - prepared = transfer._prepare_batch(batch) - monkeypatch.setattr( - transfer, - "_prepare_batch", - lambda _batch: pytest.fail("hot path must not prepare descriptors"), - ) - monkeypatch.setattr( - transfer_module.torch.cuda, - "stream", - lambda _stream: pytest.fail("must not switch streams"), - ) - transfer._copy_batch(batch, prepared) - - expected = ( - (1160, 4096, 64), - (2160, 5096, 64), - (3160, 6096, 64), - (1064, 7032, 32), - (2064, 8032, 32), - (3064, 9032, 32), - ) - assert batch_memcpy.prepared == [expected] - assert batch_memcpy.enqueued == [(expected, 123)] - assert prepared.remote_entries == {3: [(0, [remote_inbound], staging)]} - - -def test_nixl_prepare_transfer_batches_match_staging_depth(): - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer.staging_depth = 2 - transfer.staging = ["staging-0", "staging-1"] - seen_batches = [] - - def prepare_batch(batch): - seen_batches.append(batch) - return f"prepared-{len(seen_batches)}" - - transfer._prepare_batch = prepare_batch - layer_plans = [(3, "plan-3"), (4, "plan-4"), (5, "plan-5")] - - prepared_batches = transfer.prepare_transfer(layer_plans) - - assert prepared_batches == [ - (seen_batches[0], "prepared-1"), - (seen_batches[1], "prepared-2"), + TransferStep(0, 0, 0, 1), + TransferStep(1, 1, 1, 3), ] - assert [[(layer, buffer) for layer, _plan, buffer, _staging in batch] for batch in seen_batches] == [ - [(3, 0), (4, 1)], - [(5, 0)], - ] - - -def test_cuda_batch_memcpy_cuda13_abi_and_descriptor_layout(): - class Function: - def __init__(self, callback): - self.callback = callback - self.restype = None - self.argtypes = None - - def __call__(self, *args): - return self.callback(*args) - - class Library: - def __init__(self): - def get_version(pointer): - ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 13000 - return 0 - - self.cudaRuntimeGetVersion = Function(get_version) - self.cudaMemcpyBatchAsync = Function(lambda *args: self.calls.append(args) or 0) - self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") - self.calls = [] - - import ctypes - - library = Library() - batch_memcpy = _CudaBatchMemcpy(library) - prepared = batch_memcpy.prepare(((101, 201, 64), (102, 202, 128))) - batch_memcpy.enqueue(prepared, 777) - - assert len(library.cudaMemcpyBatchAsync.argtypes) == 8 - assert library.calls[0][3] == 2 - assert library.calls[0][6] == 1 - assert library.calls[0][7].value == 777 - assert [pointer for pointer in library.calls[0][0]] == [201, 202] - assert [pointer for pointer in library.calls[0][1]] == [101, 102] - assert list(library.calls[0][2]) == [64, 128] - attrs = library.calls[0][4]._obj - assert attrs.srcAccessOrder == 1 - assert attrs.srcLocHint.type == attrs.srcLocHint.id == 0 - assert attrs.dstLocHint.type == attrs.dstLocHint.id == 0 - assert attrs.flags == 1 - - -def test_cuda_batch_memcpy_rejects_unsupported_runtime_and_invalid_descriptors(): - class Function: - def __init__(self, callback): - self.callback = callback - self.restype = None - self.argtypes = None - - def __call__(self, *args): - return self.callback(*args) - - class OldRuntimeLibrary: - def __init__(self): - def get_version(pointer): - ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 12080 - return 0 - - self.cudaRuntimeGetVersion = Function(get_version) - self.cudaMemcpyBatchAsync = Function(lambda *_args: 0) - self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") - - import ctypes - with pytest.raises(RuntimeError, match="13.0"): - _CudaBatchMemcpy(OldRuntimeLibrary()) + grouped = PinnedMemoryEPLBTransfer._group_steps_by_source(plan) - class FutureRuntimeLibrary: - def __init__(self): - def get_version(pointer): - ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 14000 - return 0 - - self.cudaRuntimeGetVersion = Function(get_version) - self.cudaMemcpyBatchAsync = Function(lambda *_args: 0) - self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + assert [source for source, _steps in grouped] == [(0, 1), (1, 3)] + assert grouped[1][1] == [plan[0], plan[2]] - with pytest.raises(RuntimeError, match="13.x"): - _CudaBatchMemcpy(FutureRuntimeLibrary()) - class MissingBatchSymbolLibrary: +def test_pinned_transfer_copies_source_row_through_cpu_buffer(monkeypatch): + class Stream: def __init__(self): - def get_version(pointer): - ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 13000 - return 0 - - self.cudaRuntimeGetVersion = Function(get_version) - self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + self.synchronize_count = 0 - with pytest.raises(RuntimeError, match="cudaMemcpyBatchAsync"): - _CudaBatchMemcpy(MissingBatchSymbolLibrary()) - with pytest.raises(ValueError, match="at least one"): - _CudaBatchMemcpy.prepare(()) - with pytest.raises(ValueError, match="positive"): - _CudaBatchMemcpy.prepare(((1, 2, 0),)) - - -def test_nixl_transfer_fails_fast_without_cuda13_batch_memcpy(monkeypatch): - failure = RuntimeError("missing cudaMemcpyBatchAsync") - - def unavailable(): - raise failure + def synchronize(self): + self.synchronize_count += 1 + transfer = object.__new__(PinnedMemoryEPLBTransfer) + transfer.global_rank = 0 + transfer.transfer_group = object() + transfer._copy_stream = Stream() + transfer.live = [[("weight", torch.tensor([[1.0, 2.0], [3.0, 4.0]]))]] + transfer.pinned_rows = [("weight", torch.empty(2))] + transfer.staging = [("weight", torch.zeros((2, 2)))] + broadcasts = [] + monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: nullcontext()) monkeypatch.setattr( - transfer_module._EPLBTransferBase, - "__init__", - lambda self, *_args: setattr(self, "device", "mock"), + transfer_module.dist, + "broadcast", + lambda tensor, src, group: broadcasts.append((tensor.clone(), src, group)), ) - monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) - monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", unavailable) - with pytest.raises(RuntimeError, match="missing cudaMemcpyBatchAsync") as exc_info: - transfer_module.NixlEPLBTransfer([object()], object(), 0, 1) - assert exc_info.value is failure - - -@pytest.mark.parametrize( - ("maps", "open_raises", "expected"), - [ - ( - "7f /cuda-12/libcudart.so.12.8\n" - "7f /first/libcudart.so.13 (deleted)\n" - "7f /second/libcudart.so.13\n" - "7f /cuda-14/libcudart.so.14\n" - "7f /cuda-130/libcudart.so.130\n", - False, - "/first/libcudart.so.13", - ), - ("7f /cuda-12/libcudart.so.12\n7f /cuda-130/libcudart.so.130\n", False, None), - ("", True, None), - ], -) -def test_cuda_batch_memcpy_finds_first_loaded_cuda13_runtime(monkeypatch, maps, open_raises, expected): - def open_maps(_path): - if open_raises: - raise OSError("maps unavailable") - return io.StringIO(maps) - - monkeypatch.setattr(builtins, "open", open_maps) - assert _CudaBatchMemcpy._find_loaded_cudart() == expected - - -@pytest.mark.parametrize(("layers", "expected_depth"), [(1, 1), (8, 8), (9, 8), (43, 8)]) -def test_nixl_transfer_bounds_staging_depth(monkeypatch, layers, expected_depth): - def base_init(self, *_args): - self.device = "mock-device" - - monkeypatch.setattr(transfer_module._EPLBTransferBase, "__init__", base_init) - monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) - monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", lambda: object()) - monkeypatch.setattr(transfer_module.NixlEPLBTransfer, "_init_ipc_metadata", lambda self: None) - monkeypatch.setattr(transfer_module.NixlEPLBTransfer, "_init_push_layouts", lambda self: None) - - transfer = transfer_module.NixlEPLBTransfer([object()] * layers, object(), 0, 1) - - assert transfer.staging_depth == expected_depth - - -def test_nixl_remote_read_cache_reuses_exact_batch_key_and_releases_on_shutdown(): - class Tensor: - nbytes = 8 - - def __init__(self, pointer): - self.pointer = pointer - def data_ptr(self): - return self.pointer - - def get_device(self): - return 0 - - def __getitem__(self, _): - return self - - class Agent: - def __init__(self): - self.prepared = 0 - self.made = 0 - self.released_xfers = 0 - self.released_dlists = 0 - self.removed_agents = [] - - def get_xfer_descs(self, descriptors, _): - return descriptors - - def prep_xfer_dlist(self, *_args, **_kwargs): - self.prepared += 1 - return f"dlist-{self.prepared}" - - def make_prepped_xfer(self, *_args, **_kwargs): - self.made += 1 - return f"xfer-{self.made}" - - def query_xfer_backend(self, _): - return "UCX" - - def release_xfer_handle(self, _): - self.released_xfers += 1 - - def release_dlist_handle(self, _): - self.released_dlists += 1 - - def remove_remote_agent(self, remote_name): - self.removed_agents.append(remote_name) - - agent = Agent() - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._nixl_agent = agent - transfer._xfer_cache = {} - transfer._used_xfer_cache_keys = set() - transfer._remote_agents = {1: "remote-1"} - transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} - transfer._registered_descs = None - transfer.live = [[("weight", Tensor(100))]] - staging = [("weight", Tensor(200))] - first_entries = [(0, [TransferStep(0, 0, 1, 0)], staging)] - - first = transfer._get_remote_read(1, first_entries) - assert transfer._get_remote_read(1, first_entries) is first - assert agent.made == 1 - - changed_entries = [(0, [TransferStep(0, 1, 1, 0)], staging)] - transfer._get_remote_read(1, changed_entries) - assert agent.made == 2 - - transfer.shutdown() - assert agent.released_xfers == 2 - assert agent.released_dlists == 4 - assert agent.removed_agents == ["remote-1"] - - -def test_nixl_remote_read_cache_is_bounded_to_the_current_transfer_generation( - monkeypatch, -): - class Tensor: - nbytes = 8 - - def __init__(self, pointer): - self.pointer = pointer - - def data_ptr(self): - return self.pointer + transfer._copy_layer( + 0, + [ + TransferStep(dst_rank=0, dst_slot=1, src_rank=0, src_local_row=1), + TransferStep(dst_rank=1, dst_slot=0, src_rank=0, src_local_row=1), + ], + ) - def get_device(self): - return 0 + assert len(broadcasts) == 1 + assert torch.equal(broadcasts[0][0], torch.tensor([3.0, 4.0])) + assert broadcasts[0][1:] == (0, transfer.transfer_group) + assert torch.equal(transfer.staging[0][1][1], torch.tensor([3.0, 4.0])) + assert torch.count_nonzero(transfer.staging[0][1][0]) == 0 + assert transfer._copy_stream.synchronize_count == 2 - def __getitem__(self, _): - return self - class Agent: +def test_pinned_transfer_waits_for_each_layer_commit_before_reusing_staging(monkeypatch): + class Event: def __init__(self): - self.made = 0 - self.released_xfers = 0 - self.released_dlists = 0 - - def get_xfer_descs(self, descriptors, _): - return descriptors + self.synchronize_count = 0 - def prep_xfer_dlist(self, *_args, **_kwargs): - return object() - - def make_prepped_xfer(self, *_args, **_kwargs): - self.made += 1 - return object() - - def query_xfer_backend(self, _): - return "UCX" - - def release_xfer_handle(self, _): - self.released_xfers += 1 - - def release_dlist_handle(self, _): - self.released_dlists += 1 + def synchronize(self): + self.synchronize_count += 1 - def remove_remote_agent(self, _): + def record(self, _stream): pass - agent = Agent() - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._nixl_agent = agent - transfer._xfer_cache = {} - transfer._used_xfer_cache_keys = set() - transfer._remote_agents = {1: "remote-1"} - transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} - transfer._registered_descs = None - transfer._ipc_staging = {} - transfer.live = [[("weight", Tensor(100))]] - transfer.device = torch.device("cuda", 0) - transfer.transfer_group = object() + transfer = object.__new__(PinnedMemoryEPLBTransfer) + transfer.device = "cuda:0" transfer.global_rank = 0 - transfer.world_size = 2 - transfer.staging_depth = 1 - transfer.staging = [[]] + transfer.live = [[], []] + transfer.staging = [] transfer.num_experts_per_rank = 0 - transfer._release = [threading.Event()] - transfer._release[0].set() - transfer._consumed_events = [object()] - transfer._consumed_recorded = [False] - transfer._changed_dst_slots = [()] - transfer._pending = deque() + transfer._release = threading.Event() + transfer._release.set() + transfer._consumed_event = Event() + transfer._consumed_recorded = False + transfer._changed_dst_slots = () + transfer._pending = transfer_module.deque() transfer._pending_lock = threading.Lock() transfer._error = None transfer._thread = None - transfer._needs_staging_reuse_barrier = False - - staging = [("weight", Tensor(200))] - entries_a = [(0, [TransferStep(0, 0, 1, 0)], staging)] - entries_b = [(0, [TransferStep(0, 1, 1, 0)], staging)] - generation = [entries_a] - - def copy_batch(_batch, _prepared_batch): - transfer._get_remote_read(1, generation[0]) - - transfer._copy_batch = copy_batch + copied = [] + transfer._copy_layer = lambda layer_index, _plan: copied.append(layer_index) monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + monkeypatch.setattr(transfer_module.torch.cuda, "current_stream", lambda: object()) - prepared_batches = [([(0, [], 0, transfer.staging[0])], None)] - transfer.start([(0, [])], prepared_batches) - transfer.finish() - assert agent.made == 1 - assert agent.released_xfers == agent.released_dlists == 0 - assert len(transfer._xfer_cache) == 1 - - # The real manager releases each staging buffer through commit(). This - # focused cache test has no commits, so model that hand-off before the - # next generation reuses buffer zero. - transfer._release[0].set() - transfer.start([(0, [])], prepared_batches) - transfer.finish() - assert agent.made == 1 - assert agent.released_xfers == agent.released_dlists == 0 - assert len(transfer._xfer_cache) == 1 + transfer.start([(0, []), (1, [])]) + deadline = time.monotonic() + 2 + while not transfer.pending_layers() and time.monotonic() < deadline: + time.sleep(0.001) + assert transfer.pending_layers() == [(0, 0)] + assert copied == [0] - generation[:] = [entries_b] - transfer._release[0].set() - transfer.start([(0, [])], prepared_batches) + transfer.commit(0, 0) + deadline = time.monotonic() + 2 + while not transfer.pending_layers() and time.monotonic() < deadline: + time.sleep(0.001) + assert transfer.pending_layers() == [(1, 0)] + assert copied == [0, 1] + assert transfer._consumed_event.synchronize_count == 1 + transfer.commit(1, 0) transfer.finish() - assert agent.made == 2 - assert agent.released_xfers == 1 - assert agent.released_dlists == 2 - assert len(transfer._xfer_cache) == 1 - transfer.shutdown() - - -def test_nixl_ipc_metadata_exports_staging_per_local_target(monkeypatch): - from lightllm.server.router.model_infer.mode_backend.pd import p2p_fix - - class Tensor: - shape = (4, 2) - dtype = torch.float16 - device = torch.device("cuda", 0) - nbytes = 16 - - def __init__(self, label): - self.label = label - - def numel(self): - return 3 - - def __getitem__(self, _index): - return self - - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer.transfer_group = object() - transfer.global_rank = 0 - transfer.world_size = 3 - transfer.device = torch.device("cuda", 0) - transfer.live = [[("w13.weight", Tensor("local-w13")), ("w2.weight", Tensor("local-w2"))]] - transfer.staging_depth = 8 - transfer.staging = [[("w13.weight", Tensor(f"local-staging-{index}"))] for index in range(8)] - transfer._ipc_staging = {} - - reduce_calls, rebuild_calls, gathers = [], [], [] - - def reduce_tensor(tensor): - reduce_calls.append(tensor.label) - return None, (f"export-{tensor.label}",) - - def rebuild_tensor(export): - rebuild_calls.append(export) - return Tensor(export) - - source_one = [[("w13.weight", (4, 2), torch.float16, (f"rank1-staging-{index}",))] for index in range(8)] - - def all_gather(output, value, **_kwargs): - gathers.append(value) - if len(gathers) == 1: - output[:] = ["node-a", "node-a", "node-b"] - else: - output[:] = [value, {0: {"staging": source_one}}, {}] - - monkeypatch.setattr(p2p_fix, "reduce_tensor", reduce_tensor) - monkeypatch.setattr(p2p_fix, "p2p_fix_rebuild_cuda_tensor", rebuild_tensor) - monkeypatch.setattr(transfer_module.dist, "all_gather_object", all_gather) - monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) - - transfer._init_ipc_metadata() - - assert reduce_calls == [*(f"local-staging-{index}" for index in range(8))] - assert rebuild_calls == [*(f"rank1-staging-{index}" for index in range(8))] - assert transfer._same_node_ranks == {0, 1} - assert transfer._cross_node_ranks == {2} - assert transfer._needs_staging_reuse_barrier - assert [name for name, _ in transfer._ipc_staging[1][0]] == ["w13.weight"] - - -def test_nixl_copy_batch_source_pushes_local_rows_and_keeps_remote_ucx_reads( - monkeypatch, -): - class Stream: - def __init__(self): - self.synchronized = 0 - self.cuda_stream = 123 - - def synchronize(self): - self.synchronized += 1 - - stream = Stream() - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._push_stream = stream - transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) - remote_reads, waited_xfers, enqueued = [], [], [] - transfer._get_remote_read = lambda src_rank, entries: remote_reads.append((src_rank, entries)) or ( - None, - None, - "xfer", - ) - transfer._wait_xfers = waited_xfers.extend - - monkeypatch.setattr( - transfer_module.torch.cuda, - "stream", - lambda _stream: pytest.fail("must not switch streams"), - ) - prepared_push = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") - transfer._copy_batch([], prepared_push) - assert enqueued == [("push", stream.cuda_stream)] - - remote_step = TransferStep(0, 1, 2, 0) - prepared_remote = transfer_module.NixlEPLBTransfer._PreparedBatch({2: [(0, [remote_step], [])]}, None) - transfer._copy_batch([], prepared_remote) - assert [rank for rank, _ in remote_reads] == [2] - assert waited_xfers == [(None, None, "xfer")] - - -def test_nixl_copy_batch_self_only_rank_pushes_and_synchronizes(monkeypatch): - class Stream: - def __init__(self): - self.synchronized = 0 - self.cuda_stream = 456 - - def synchronize(self): - self.synchronized += 1 - - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._push_stream = Stream() - enqueued = [] - transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) - transfer._wait_xfers = lambda _xfers: None - monkeypatch.setattr( - transfer_module.torch.cuda, - "stream", - lambda _stream: pytest.fail("must not switch streams"), - ) - - prepared = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") - transfer._copy_batch([], prepared) - - assert enqueued == [("push", transfer._push_stream.cuda_stream)] - assert transfer._push_stream.synchronized == 1 -def test_manager_constructs_nixl_transfer(monkeypatch): +def test_manager_constructs_pinned_memory_transfer(monkeypatch): weight = type( "Weight", (), @@ -3504,7 +2841,7 @@ def new_group(*args, **kwargs): transfer_calls = [] monkeypatch.setattr( manager_module, - "NixlEPLBTransfer", + "PinnedMemoryEPLBTransfer", lambda weights, group, rank, world_size: ( transfer_calls.append((weights, group, rank, world_size)) or transfer ), diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index b334a6fa15..5893cda007 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -1,9 +1,7 @@ -"""NIXL EPLB correctness tests and a two-GPU 512 MiB micro-performance test.""" +"""Multi-GPU correctness test for pinned-memory EPLB transfers.""" import os -import random import socket -import statistics import time from types import SimpleNamespace @@ -13,26 +11,9 @@ import torch.multiprocessing as mp from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( - NixlEPLBTransfer, - align_target_placement, + PinnedMemoryEPLBTransfer, build_transfer_plan, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_initial_local_expert_ids, -) - -pytest.importorskip("nixl", reason="NIXL package is required") - - -def _initial_extra_expert_placement(num_logical_experts, world_size, num_redundant_experts_per_rank): - num_primary_experts_per_rank = num_logical_experts // world_size - initial_local_expert_ids_by_rank = build_initial_local_expert_ids( - num_logical_experts, world_size, num_redundant_experts_per_rank - ) - return torch.tensor( - [expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids_by_rank], - dtype=torch.int64, - ) class _Pack: @@ -42,6 +23,27 @@ def __init__(self, weight, weight_scale): self.weight_zero_point = None +class _FakeWeight: + def __init__(self, rank, layer_index): + self.fuse_moe_impl = SimpleNamespace( + num_primary_experts_per_rank=2, + num_redundant_experts_per_rank=1, + ) + logical_ids = ([0, 1, 2], [2, 3, 0])[rank] + self.w13 = self._pack(logical_ids, layer_index, 0) + self.w2 = self._pack(logical_ids, layer_index, 10) + + @staticmethod + def _pack(logical_ids, layer_index, offset): + values = torch.tensor( + [[layer_index * 100 + expert + offset] for expert in logical_ids], + dtype=torch.float16, + device="cuda", + ) + scales = values.to(torch.float32) + 0.5 + return _Pack(values, scales) + + def _free_port(): sock = socket.socket() sock.bind(("127.0.0.1", 0)) @@ -50,411 +52,60 @@ def _free_port(): return port -class _FakeWeight: - def __init__(self, rank, layer_index, row_elements): - self.n_routed_experts = 32 - self.fuse_moe_impl = SimpleNamespace( - n_routed_experts=32, - num_primary_experts_per_rank=16, - num_redundant_experts_per_rank=16, - logical_to_physical_map=torch.cat( - ( - torch.ones((32, 1), dtype=torch.int32, device="cuda"), - torch.zeros((32, 3), dtype=torch.int32, device="cuda"), - ), - dim=1, - ), - route_counter=torch.zeros((32,), dtype=torch.int64, device="cuda"), - ) - base = rank * 100 + layer_index * 100 - self.w13 = self._pack(base, row_elements) - self.w2 = self._pack(base + 10, row_elements) - - @staticmethod - def _pack(base, row_elements): - weight = torch.empty((32, row_elements), dtype=torch.float16, device="cuda") - for row in range(weight.shape[0]): - weight[row].fill_(base + row) - scale = torch.empty((32, 1), dtype=torch.float32, device="cuda") - for row in range(scale.shape[0]): - scale[row].fill_(base + row + 0.5) - return _Pack(weight, scale) - - -def _wait_for_ready_prefix(transfer, control_group): +def _wait_for_ready_layer(transfer, control_group): deadline = time.monotonic() + 30 - while True: + while time.monotonic() < deadline: pending = transfer.pending_layers() ready_count = torch.tensor([len(pending)], dtype=torch.int32) dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=control_group) - if int(ready_count.item()) > 0: - return pending[: int(ready_count.item())] - if time.monotonic() >= deadline: - raise TimeoutError("EPLB transfer worker did not publish a globally ready layer") + if int(ready_count.item()) == 1: + return pending[0] time.sleep(0.001) + raise TimeoutError("EPLB transfer worker did not publish a globally ready layer") -def _run_layers( - transfer, - control_group, - layer_plans, - callback=lambda layer_index: None, - before_commit_callback=lambda layer_index: None, -): - transfer.start(layer_plans, transfer.prepare_transfer(layer_plans)) - committed = 0 - while committed < len(layer_plans): - pending = _wait_for_ready_prefix(transfer, control_group) - assert len(pending) <= len(layer_plans) - committed - for layer_index, buffer_index in pending: - assert layer_index == layer_plans[committed][0] - before_commit_callback(layer_index) - transfer.commit( - layer_index, - buffer_index, - lambda layer_index=layer_index: callback(layer_index), - ) - committed += 1 - transfer.finish() - - -def _assert_correctness(weights, rank, source_rows_by_dst_slot): - if rank == 0: - for layer_index in (0, len(weights) - 1): - base = 100 + layer_index * 100 - for dst_slot, src_row in enumerate(source_rows_by_dst_slot): - dst_row = 16 + dst_slot - assert torch.all(weights[layer_index].w13.weight[dst_row] == base + src_row) - assert torch.all(weights[layer_index].w13.weight_scale[dst_row] == base + src_row + 0.5) - assert torch.all(weights[layer_index].w2.weight[dst_row] == base + src_row + 10) - assert torch.all(weights[layer_index].w2.weight_scale[dst_row] == base + src_row + 10.5) - - -def _benchmark(transfer, control_group, layer_plans, payload): - for _ in range(3): - _run_layers(transfer, control_group, layer_plans) - dist.barrier(group=control_group) - started = time.perf_counter() - for _ in range(8): - _run_layers(transfer, control_group, layer_plans) - torch.cuda.synchronize() - return payload * 8 / (time.perf_counter() - started) / 1e9 - - -def _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload): - batch = [ - (layer_index, plan, buffer_index, transfer.staging[buffer_index]) - for buffer_index, (layer_index, plan) in enumerate(layer_plans) - ] - # Measure the precompiled hot path only; planning and descriptor construction are off the timer. - prepared_batch = transfer._prepare_batch(batch) - for _ in range(3): - transfer._copy_batch(batch, prepared_batch) - dist.barrier(group=control_group) - samples = [] - for _ in range(20): - started = time.perf_counter() - transfer._copy_batch(batch, prepared_batch) - samples.append(payload / (time.perf_counter() - started) / 1e9) - torch.cuda.synchronize() - median = statistics.median(samples) - print( - f"NIXL _copy_batch payload={payload / 2**20:.1f} MiB; " - f"min={min(samples):.2f} GB/s median={median:.2f} " - f"mean={statistics.mean(samples):.2f} max={max(samples):.2f}", - flush=True, - ) - return median - - -def _eplb_worker(rank, port, queue): +def _worker(rank, port): os.environ["MASTER_ADDR"] = "127.0.0.1" os.environ["MASTER_PORT"] = str(port) torch.cuda.set_device(rank) dist.init_process_group("gloo", rank=rank, world_size=2) control_group = dist.new_group([0, 1], backend="gloo") transfer_group = dist.new_group([0, 1], backend="gloo") - # Eight layers × 16 changed experts × two 2 MiB rows = 512 MiB useful remote weight payload. - row_elements = int(os.getenv("LIGHTLLM_EPLB_TEST_ROW_ELEMENTS", str(1024 * 1024))) - current = torch.tensor([list(range(16)), list(range(16))]) - source_rows_by_dst_slot = list(range(16)) - random.Random(20260731).shuffle(source_rows_by_dst_slot) - # Rank 0 receives the reverse logical-expert range in a deterministic random slot order. - # Consequently every descriptor has a distinct source and destination row. - target = torch.tensor([[16 + source_row for source_row in source_rows_by_dst_slot], list(range(16))]) - plan = build_transfer_plan(current, target, 32, 2, 2) - assert [step.src_local_row for step in plan if step.dst_rank == 0] == source_rows_by_dst_slot - benchmark_layer_count = 8 - layer_count = benchmark_layer_count + 1 - row_payload = ( - 2 * row_elements * torch.empty((), dtype=torch.float16).element_size() - + 2 * torch.empty((), dtype=torch.float32).element_size() - ) - payload = benchmark_layer_count * 16 * row_payload - weights = [_FakeWeight(rank, layer_index, row_elements) for layer_index in range(layer_count)] - transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=2) - assert transfer.staging_depth == 8 - assert transfer._eplb_impls[0] is weights[0].fuse_moe_impl - assert all(tensor.is_cuda for staging in transfer.staging for _, tensor in staging) - - wrap_layer_plans = [(layer_index, plan) for layer_index in range(layer_count)] - - def delay_rank_zero_first_commit(layer_index): - if rank == 0 and layer_index == 0: - time.sleep(0.1) - - _run_layers( - transfer, - control_group, - wrap_layer_plans, - before_commit_callback=delay_rank_zero_first_commit, - ) - torch.cuda.synchronize() - _assert_correctness(weights, rank, source_rows_by_dst_slot) - layer_plans = wrap_layer_plans[:benchmark_layer_count] - nixl_copy_batch = _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload) - nixl_bandwidth = _benchmark(transfer, control_group, layer_plans, payload) - transfer.shutdown() - - gathered = [None, None] - dist.all_gather_object(gathered, (nixl_bandwidth, nixl_copy_batch), group=control_group) - if rank == 0: - queue.put((payload, *gathered[0])) - dist.barrier(group=control_group) - dist.destroy_process_group() - - -@pytest.mark.skipif( - not torch.cuda.is_available() or torch.cuda.device_count() < 2, - reason="requires two CUDA GPUs", -) -def test_eplb_transfer_two_gpu_correctness_and_microperf(): - queue = mp.get_context("spawn").SimpleQueue() - mp.spawn(_eplb_worker, args=(_free_port(), queue), nprocs=2, join=True) - payload, nixl_gbps, nixl_copy_batch_gbps = queue.get() - print( - f"EPLB remote payload/round: {payload / 2**20:.1f} MiB; " - f"NIXL={nixl_gbps:.2f} GB/s NIXL _copy_batch={nixl_copy_batch_gbps:.2f} GB/s" - ) - assert nixl_gbps > 0 - - -def _depth_value(expert, layer_index, offset): - return expert // 32 * 100 + layer_index * 5 + expert % 32 + offset - - -class _DepthWeight: - def __init__(self, rank, layer_index, initial_placement): - self.n_routed_experts = 256 - self.fuse_moe_impl = SimpleNamespace( - n_routed_experts=256, - num_primary_experts_per_rank=32, - num_redundant_experts_per_rank=4, - logical_to_physical_map=torch.cat( - ( - torch.ones((256, 1), dtype=torch.int32, device="cuda"), - torch.zeros((256, 9), dtype=torch.int32, device="cuda"), - ), - dim=1, - ), - route_counter=torch.zeros((256,), dtype=torch.int64, device="cuda"), - ) - logical_ids = list(range(rank * 32, (rank + 1) * 32)) + initial_placement[rank].tolist() - self.w13 = self._pack(logical_ids, layer_index, 0) - self.w2 = self._pack(logical_ids, layer_index, 2) - - @staticmethod - def _pack(logical_ids, layer_index, offset): - weight = torch.empty((36, 64), dtype=torch.float16, device="cuda") - scale = torch.empty((36, 1), dtype=torch.float32, device="cuda") - for row, expert in enumerate(logical_ids): - value = _depth_value(expert, layer_index, offset) - weight[row].fill_(value) - scale[row].fill_(value + 0.25) - return _Pack(weight, scale) - - -def _depth_target(layer_index): - return torch.tensor([[((dst + layer_index + slot + 1) % 8) * 32 + slot for slot in range(4)] for dst in range(8)]) - - -def _wait_all_pending(transfer, group, expected_count): - deadline = time.monotonic() + 30 - while True: - pending = transfer.pending_layers() - ready_count = torch.tensor([len(pending)], dtype=torch.int32) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=group) - if int(ready_count.item()) == expected_count: - return pending - if time.monotonic() > deadline: - raise TimeoutError(f"expected {expected_count} pending layers, got {pending}") - time.sleep(0.001) - - -def _clone_depth_live(weights): - return [ - [tensor.detach().clone() for _, tensor in transfer_tensors] - for transfer_tensors in [ - [ - ("w13.weight", weight.w13.weight), - ("w13.scale", weight.w13.weight_scale), - ("w2.weight", weight.w2.weight), - ("w2.scale", weight.w2.weight_scale), - ] - for weight in weights - ] - ] - -def _assert_depth_snapshot(weights, snapshot, layer_indices=None, primary_only=False): - if layer_indices is None: - layer_indices = range(len(weights)) - for layer_index in layer_indices: - live_tensors = ( - weights[layer_index].w13.weight, - weights[layer_index].w13.weight_scale, - weights[layer_index].w2.weight, - weights[layer_index].w2.weight_scale, - ) - for live, expected in zip(live_tensors, snapshot[layer_index]): - if primary_only: - live = live[:32] - expected = expected[:32] - torch.testing.assert_close(live, expected) - - -def _assert_depth_staging(rank, layer_plans, target_placements, pending, transfer): - for buffer_index, ((layer_index, plan), pending_item) in enumerate(zip(layer_plans, pending)): - assert pending_item == (layer_index, buffer_index) - expected = {step.dst_slot: step for step in plan if step.dst_rank == rank} - for dst_slot, step in expected.items(): - expert = int(target_placements[layer_index][rank, dst_slot]) - base = _depth_value(expert, layer_index, 0) - staging = transfer.staging[buffer_index] - assert torch.all(staging[0][1][dst_slot] == base) - assert torch.all(staging[1][1][dst_slot] == base + 0.25) - assert torch.all(staging[2][1][dst_slot] == base + 2) - assert torch.all(staging[3][1][dst_slot] == base + 2.25) - - -def _assert_depth_live(weights, rank, layer_plans, target_placements): - for layer_index, plan in layer_plans: - for step in plan: - if step.dst_rank != rank: - continue - expert = int(target_placements[layer_index][rank, step.dst_slot]) - base = _depth_value(expert, layer_index, 0) - assert torch.all(weights[layer_index].w13.weight[32 + step.dst_slot] == base) - assert torch.all(weights[layer_index].w13.weight_scale[32 + step.dst_slot] == base + 0.25) - assert torch.all(weights[layer_index].w2.weight[32 + step.dst_slot] == base + 2) - assert torch.all(weights[layer_index].w2.weight_scale[32 + step.dst_slot] == base + 2.25) - - -def _assert_peer_coverage(layer_plans, require_redundant_source=False): - steps = [step for _, plan in layer_plans for step in plan] - assert {step.dst_rank for step in steps} == set(range(8)) - assert {step.src_rank for step in steps} == set(range(8)) - for source_rank in range(8): - assert len({step.dst_rank for step in steps if step.src_rank == source_rank}) > 1 - if require_redundant_source: - assert any(step.src_local_row >= 32 for step in steps) - - -def _depth_worker(rank, port): - os.environ["MASTER_ADDR"] = "127.0.0.1" - os.environ["MASTER_PORT"] = str(port) - torch.cuda.set_device(rank) - dist.init_process_group("gloo", rank=rank, world_size=8) - control_group = dist.new_group(list(range(8)), backend="gloo") - transfer_group = dist.new_group(list(range(8)), backend="gloo") - initial_placement = _initial_extra_expert_placement(256, 8, 4) - weights = [_DepthWeight(rank, layer_index, initial_placement) for layer_index in range(9)] - transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=8) - assert transfer.staging_depth == 8 - staging_bytes = sum(tensor.nbytes for staging in transfer.staging for _, tensor in staging) - one_layer_staging_bytes = sum(tensor[:4].nbytes for _, tensor in transfer.live[0]) - assert staging_bytes == 8 * one_layer_staging_bytes - current = initial_placement - first_order = [8, 0, 5, 1, 7, 3, 6, 2, 4] - first_targets = { - layer_index: align_target_placement(current, _depth_target(layer_index)) for layer_index in range(9) - } - first_plans = [ - ( - layer_index, - build_transfer_plan(current, first_targets[layer_index], 256, 8, 8), - ) - for layer_index in first_order - ] - _assert_peer_coverage(first_plans) - first_snapshot = _clone_depth_live(weights) - transfer.start(first_plans, transfer.prepare_transfer(first_plans)) - committed = 0 - pending = _wait_all_pending(transfer, control_group, 8) - _assert_depth_staging(rank, first_plans[:8], first_targets, pending, transfer) - _assert_depth_snapshot(weights, first_snapshot) - dist.barrier(group=control_group) - if rank == 0: - time.sleep(0.1) - for layer_index, buffer_index in pending: - _assert_depth_snapshot(weights, first_snapshot, [layer for layer, _ in first_plans[committed:]]) - transfer.commit(layer_index, buffer_index) - committed += 1 - pending = _wait_all_pending(transfer, control_group, 1) - _assert_depth_staging(rank, first_plans[8:], first_targets, pending, transfer) - _assert_depth_snapshot(weights, first_snapshot, [first_plans[8][0]]) - dist.barrier(group=control_group) - if rank == 0: - time.sleep(0.1) - for layer_index, buffer_index in pending: - _assert_depth_snapshot(weights, first_snapshot, [layer for layer, _ in first_plans[committed:]]) + weights = [_FakeWeight(rank, layer_index) for layer_index in range(2)] + transfer = PinnedMemoryEPLBTransfer(weights, transfer_group, rank, world_size=2) + assert all(row.is_pinned() for _, row in transfer.pinned_rows) + current = torch.tensor([[2], [0]]) + target = torch.tensor([[3], [1]]) + plan = build_transfer_plan(current, target, num_logical_experts=4, world_size=2, node_world_size=2) + layer_plans = [(0, plan), (1, plan)] + transfer.start(layer_plans) + + for expected_layer in range(2): + layer_index, buffer_index = _wait_for_ready_layer(transfer, control_group) + assert layer_index == expected_layer + expected_expert = 3 if rank == 0 else 1 + expected_w13 = expected_layer * 100 + expected_expert + expected_w2 = expected_w13 + 10 + assert torch.all(transfer.staging[0][1][0] == expected_w13) + assert torch.all(transfer.staging[1][1][0] == expected_w13 + 0.5) + assert torch.all(transfer.staging[2][1][0] == expected_w2) + assert torch.all(transfer.staging[3][1][0] == expected_w2 + 0.5) transfer.commit(layer_index, buffer_index) - committed += 1 - transfer.finish() - torch.cuda.synchronize() - dist.barrier(group=control_group) - _assert_depth_live(weights, rank, first_plans, first_targets) - second_order = [7, 2, 4] - second_targets = { - layer_index: align_target_placement( - first_targets[layer_index], - torch.tensor([[((dst + layer_index + slot + 3) % 8) * 32 + slot for slot in range(4)] for dst in range(8)]), - ) - for layer_index in second_order - } - second_plans = [ - ( - layer_index, - build_transfer_plan(first_targets[layer_index], second_targets[layer_index], 256, 8, 8), - ) - for layer_index in second_order - ] - _assert_peer_coverage(second_plans, require_redundant_source=True) - second_snapshot = _clone_depth_live(weights) - transfer.start(second_plans, transfer.prepare_transfer(second_plans)) - pending = _wait_all_pending(transfer, control_group, len(second_plans)) - _assert_depth_staging(rank, second_plans, second_targets, pending, transfer) - _assert_depth_snapshot(weights, second_snapshot) - dist.barrier(group=control_group) - if rank == 0: - time.sleep(0.1) - for layer_index, buffer_index in pending: - transfer.commit(layer_index, buffer_index) transfer.finish() torch.cuda.synchronize() - dist.barrier(group=control_group) - _assert_depth_live(weights, rank, second_plans, second_targets) - _assert_depth_snapshot(weights, second_snapshot, set(range(9)) - set(second_order)) - _assert_depth_snapshot(weights, second_snapshot, primary_only=True) - transfer.shutdown() - dist.barrier(group=control_group) + for layer_index, weight in enumerate(weights): + expected_expert = 3 if rank == 0 else 1 + expected_w13 = layer_index * 100 + expected_expert + assert torch.all(weight.w13.weight[2] == expected_w13) + assert torch.all(weight.w2.weight[2] == expected_w13 + 10) dist.destroy_process_group() @pytest.mark.skipif( - not torch.cuda.is_available() or torch.cuda.device_count() < 8, - reason="requires eight CUDA GPUs", + not torch.cuda.is_available() or torch.cuda.device_count() < 2, + reason="requires two CUDA GPUs", ) -def test_eplb_transfer_eight_gpu_bounded_staging_reuse(): - mp.spawn(_depth_worker, args=(_free_port(),), nprocs=8, join=True) +def test_eplb_pinned_memory_transfer_two_gpu_correctness(): + mp.spawn(_worker, args=(_free_port(),), nprocs=2, join=True) From 495e9b7ac6ded14974e211074c4b5c06badfdd2b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 01:52:15 +0000 Subject: [PATCH 21/72] refactor(eplb): move lifecycle control into mode backend --- lightllm/common/basemodel/basemodel.py | 13 +- .../model_infer/mode_backend/base_backend.py | 13 +- .../mode_backend/chunked_prefill/impl.py | 8 +- .../mode_backend/dp_backend/impl.py | 8 +- .../model_infer/mode_backend/eplb_manager.py | 30 ++-- .../model_infer/mode_backend/eplb_transfer.py | 45 +++--- unit_tests/common/fused_moe/test_eplb.py | 141 +++++++++--------- .../fused_moe/test_eplb_transfer_gpu.py | 12 +- 8 files changed, 128 insertions(+), 142 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 5fb8d60d0b..f59535cb13 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -70,7 +70,6 @@ class TpPartBaseModel: def __init__(self, kvargs): self.args = get_env_start_args() - self.eplb_manager = None self.run_mode = kvargs["run_mode"] self.weight_dir_ = kvargs["weight_dir"] self.max_total_token_num = kvargs["max_total_token_num"] @@ -317,14 +316,9 @@ def forward(self, model_input: ModelInput): assert model_input.mem_indexes.is_cuda if model_input.is_prefill: - model_output = self._prefill(model_input=model_input) - self._after_prefill() - return model_output - return self._decode(model_input) - - def _after_prefill(self): - if self.eplb_manager is not None: - self.eplb_manager.step() + return self._prefill(model_input=model_input) + else: + return self._decode(model_input) def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() @@ -820,7 +814,6 @@ def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input dist_group_manager.clear_deepep_buffer() model_output0.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event model_output1.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event - self._after_prefill() return model_output0, model_output1 @torch.no_grad() diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 9da7cec8fb..23452426dc 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -70,6 +70,7 @@ def __init__(self) -> None: self.enable_decode_microbatch_overlap = get_env_start_args().enable_decode_microbatch_overlap self.enable_prefill_microbatch_overlap = get_env_start_args().enable_prefill_microbatch_overlap self.spec_engine = None + self.eplb_manager = None # 控制 _get_classed_reqs 分类的参数变量,不同的 backend 具有可能需要不同的分类运行条件。 self.classed_req_no_decode = False @@ -259,8 +260,7 @@ def init_model(self, kvargs): if self.args.eplb_num_redundant_experts_per_rank > 0: from lightllm.server.router.model_infer.mode_backend.eplb_manager import EPLBManager - self.model.eplb_manager = EPLBManager(self.model) - dist.barrier() + self.eplb_manager = EPLBManager(self.model) # 启动infer_loop_thread, 启动两个线程进行推理,对于具备双batch推理折叠得场景 # 可以降低 cpu overhead,大幅提升gpu得使用率。 @@ -303,6 +303,15 @@ def infer_loop(self): def prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): raise NotImplementedError() + def _run_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): + self.prefill(event_pack=event_pack, prefill_reqs=prefill_reqs) + if self.eplb_manager is not None: + self.eplb_manager.step() + + def _poll_eplb(self): + if self.eplb_manager is not None: + self.eplb_manager.poll() + def decode(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): raise NotImplementedError() diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 12b5195551..085c66d618 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -61,10 +61,8 @@ def infer_loop(self): event_pack.wait_to_forward() - # EPLB polling performs collectives and may commit weights, so keep all ranks ordered before normal - # collectives/forward. - if self.model.eplb_manager is not None: - self.model.eplb_manager.poll() + # Keep EPLB collectives ordered before normal collectives/forward. + self._poll_eplb() self._try_read_new_reqs() @@ -80,7 +78,7 @@ def infer_loop(self): # 进行一次流同步,保证 _try_read_new_reqs 中的一些算子操作,必然已经完成。 # 防止后续的推理流程读取到显存中可能存在错误的数据。 g_infer_context.get_overlap_stream().wait_stream(torch.cuda.current_stream()) - self.prefill( + self._run_prefill( event_pack=event_pack, prefill_reqs=prefill_reqs, ) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 3070645b7e..793b729f82 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -122,10 +122,8 @@ def infer_loop(self): event_pack.wait_to_forward() - # EPLB polling performs collectives and may commit weights, so keep all ranks ordered before normal - # collectives/forward. - if self.model.eplb_manager is not None: - self.model.eplb_manager.poll() + # Keep EPLB collectives ordered before normal collectives/forward. + self._poll_eplb() self._try_read_new_reqs() @@ -150,7 +148,7 @@ def infer_loop(self): # 进行一次流同步,保证 _try_read_new_reqs 中的一些算子操作,必然已经完成。 # 防止后续的推理流程读取到显存中可能存在错误的数据。 g_infer_context.get_overlap_stream().wait_stream(torch.cuda.current_stream()) - self.prefill( + self._run_prefill( event_pack=event_pack, prefill_reqs=prefill_reqs, ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index d76c2a6106..5e57df508e 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -102,7 +102,7 @@ def __init__(self, model: TpPartBaseModel): # This control-group scalar is only touched from the main inference # thread, never by the background evaluation thread. self._control_ready_count = torch.empty(1, dtype=torch.int32) - self.transfer = PinnedMemoryEPLBTransfer(self.weights, self.transfer_group, self.global_rank, self.world_size) + self.transfer = PinnedMemoryEPLBTransfer(self.weights, self.transfer_group, self.global_rank) if self.global_rank == 0: logger.info( "eplb enabled " @@ -227,11 +227,13 @@ def _finish_rebalance(self): def _poll_in_flight(self): local_error = None try: - pending = self.transfer.pending_layers() + ready_layer = self.transfer.ready_layer() except BaseException as exc: - pending = [] + ready_layer = None local_error = exc - ready_count = self._control_count(EPLB_CONTROL_ERROR if local_error is not None else len(pending)) + ready_count = self._control_count( + EPLB_CONTROL_ERROR if local_error is not None else int(ready_layer is not None) + ) dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) ready_count = int(ready_count.item()) if ready_count < 0: @@ -240,24 +242,18 @@ def _poll_in_flight(self): raise RuntimeError("EPLB transfer worker failed on another rank") if ready_count == 0: return - if ready_count > len(pending) or ready_count > len(self.in_flight_layers): - raise RuntimeError("EPLB global ready count exceeds the local ordered prefix") + if ready_layer != self.in_flight_layers[0]: + raise RuntimeError(f"EPLB ready layer {ready_layer} does not match expected {self.in_flight_layers[0]}") from lightllm.server.router.model_infer.infer_batch import g_infer_context # Previous forward is queued on the shared overlap stream; order the # live-weight commit after it. The subsequent wait orders the next forward. torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - for layer_index, buffer_index in pending[:ready_count]: - if layer_index != self.in_flight_layers[0]: - raise RuntimeError( - f"EPLB pending layer {layer_index} does not match expected {self.in_flight_layers[0]}" - ) - self.transfer.commit( - layer_index, - buffer_index, - lambda: self._commit_layer_metadata(layer_index), - ) - self.in_flight_layers.pop(0) + self.transfer.commit( + ready_layer, + lambda: self._commit_layer_metadata(ready_layer), + ) + self.in_flight_layers.pop(0) if not self.in_flight_layers: self.transfer.finish() self._finish_rebalance() diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index 32c6d7e9bc..b6c4f69c74 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -1,7 +1,7 @@ """Layer-by-layer expert-row migration for EPLB.""" import threading -from collections import Counter, defaultdict, deque +from collections import Counter, defaultdict from dataclasses import dataclass from typing import List, Sequence, Tuple @@ -94,13 +94,10 @@ class PinnedMemoryEPLBTransfer: the staging rows and routing metadata together at a safe forward boundary. """ - backend = "pin-memory" - - def __init__(self, weights, transfer_group, global_rank, world_size): + def __init__(self, weights, transfer_group, global_rank): self._eplb_impls = [weight.fuse_moe_impl for weight in weights] self.transfer_group = transfer_group self.global_rank = global_rank - self.world_size = world_size self.num_experts_per_rank = self._eplb_impls[0].num_primary_experts_per_rank self.device = weights[0].w13.weight.device self.live = [extract_eplb_expert_tensors(weight) for weight in weights] @@ -135,9 +132,8 @@ def __init__(self, weights, transfer_group, global_rank, world_size): self._release.set() self._consumed_event = torch.cuda.Event() self._consumed_recorded = False - self._changed_dst_slots = () - self._pending = deque() - self._pending_lock = threading.Lock() + self._ready = None + self._ready_lock = threading.Lock() self._error = None self._thread = None @@ -180,8 +176,9 @@ def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]) -> No if self._thread is not None: raise RuntimeError("EPLB transfer has not been finished") self._error = None - with self._pending_lock: - self._pending.clear() + with self._ready_lock: + if self._ready is not None: + raise RuntimeError("EPLB ready layer has not been committed") def worker() -> None: try: @@ -191,32 +188,30 @@ def worker() -> None: self._release.clear() if self._consumed_recorded: self._consumed_event.synchronize() - self._changed_dst_slots = tuple( + changed_dst_slots = tuple( sorted({step.dst_slot for step in plan if step.dst_rank == self.global_rank}) ) self._copy_layer(layer_index, plan) - with self._pending_lock: - self._pending.append((layer_index, 0)) + with self._ready_lock: + self._ready = (layer_index, changed_dst_slots) except BaseException as exc: self._error = exc self._thread = threading.Thread(target=worker, name="eplb-pin-memory", daemon=True) self._thread.start() - def pending_layers(self): + def ready_layer(self): if self._error is not None: raise RuntimeError("EPLB migration worker failed") from self._error - with self._pending_lock: - return list(self._pending) - - def commit(self, layer_index: int, buffer_index: int, post_copy=None) -> None: - if buffer_index != 0: - raise RuntimeError("Pinned-memory EPLB uses a single staging buffer") - with self._pending_lock: - if not self._pending or self._pending[0] != (layer_index, buffer_index): - raise RuntimeError("EPLB commit does not match the pending FIFO") - self._pending.popleft() - changed_dst_slots = self._changed_dst_slots + with self._ready_lock: + return None if self._ready is None else self._ready[0] + + def commit(self, layer_index: int, post_copy=None) -> None: + with self._ready_lock: + if self._ready is None or self._ready[0] != layer_index: + raise RuntimeError("EPLB commit does not match the ready layer") + _, changed_dst_slots = self._ready + self._ready = None for (_, live), (_, staging) in zip(self.live[layer_index], self.staging): _commit_staging_rows(live, staging, self.num_experts_per_rank, changed_dst_slots) if post_copy is not None: diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 131c853f89..c107590e2a 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -23,6 +23,7 @@ from lightllm.server.router.model_infer.mode_backend import ( eplb_transfer as transfer_module, ) +from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( deepgemm_impl as deepgemm_module, ) @@ -2076,21 +2077,20 @@ def narrow(self, _dim, start, length): assert copies == [("live", 11, 3, "staging", 1, 3)] -def test_manager_inflight_ready_gate_commits_ordered_prefix_and_propagates_worker_error( - monkeypatch, -): +def test_manager_inflight_ready_gate_commits_one_layer_and_propagates_worker_error(monkeypatch): class Transfer: def __init__(self): - self.pending = [(0, 0), (1, 1), (2, 2)] + self.ready = 0 self.commits = [] self.finished = 0 - def pending_layers(self): - return self.pending + def ready_layer(self): + return self.ready - def commit(self, layer, buffer_index, post_copy=None): - assert self.pending.pop(0) == (layer, buffer_index) - self.commits.append((layer, buffer_index)) + def commit(self, layer, post_copy=None): + assert self.ready == layer + self.ready = None + self.commits.append(layer) if post_copy is not None: post_copy() @@ -2099,78 +2099,50 @@ def finish(self): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.transfer = Transfer() - manager.in_flight = True - manager.world_size = 2 manager.control_group = object() manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager.in_flight_layers = [0, 1, 2] - manager._commit_layer_metadata = lambda layer: committed.append(layer) - manager._finish_rebalance = lambda: finished.append(True) + manager.in_flight_layers = [0, 1] committed, finished = [], [] + manager._commit_layer_metadata = committed.append + manager._finish_rebalance = lambda: finished.append(True) operations = [] - current_stream_calls = [] class CurrentStream: def wait_stream(self, stream): operations.append(("wait", stream)) overlap_stream = object() - - def current_stream(): - current_stream_calls.append(True) - return CurrentStream() - - monkeypatch.setattr(manager_module.torch.cuda, "current_stream", current_stream) + monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: CurrentStream()) monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) - original_commit = manager.transfer.commit - - def record_commit(*args, **kwargs): - operations.append(("commit", args[0])) - return original_commit(*args, **kwargs) - - manager.transfer.commit = record_commit - def set_global_ready(count): - return lambda tensor, **kwargs: tensor.fill_(count) + return lambda tensor, **_kwargs: tensor.fill_(count) monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(0)) manager._poll_in_flight() assert manager.transfer.commits == [] - assert operations == [] - assert current_stream_calls == [] - # Local rank has three prefetched layers, but global MIN-ready only permits two. - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(2)) + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) manager._poll_in_flight() - assert manager.transfer.commits == [(0, 0), (1, 1)] - assert committed == [0, 1] + assert manager.transfer.commits == [0] + assert committed == [0] + assert operations == [("wait", overlap_stream)] assert not finished - assert manager.transfer.finished == 0 - assert operations == [("wait", overlap_stream), ("commit", 0), ("commit", 1)] - assert current_stream_calls == [True] - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) + manager.transfer.ready = 1 manager._poll_in_flight() + assert manager.transfer.commits == [0, 1] + assert committed == [0, 1] assert finished == [True] assert manager.transfer.finished == 1 - assert operations == [ - ("wait", overlap_stream), - ("commit", 0), - ("commit", 1), - ("wait", overlap_stream), - ("commit", 2), - ] - assert current_stream_calls == [True, True] - manager.in_flight_layers = [3] - manager.transfer.pending = [(9, 0)] - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) + manager.in_flight_layers = [2] + manager.transfer.ready = 9 with pytest.raises(RuntimeError, match="does not match expected"): manager._poll_in_flight() class BrokenTransfer: - def pending_layers(self): + def ready_layer(self): raise RuntimeError("boom") manager.transfer = BrokenTransfer() @@ -2192,8 +2164,8 @@ class Transfer: def __init__(self): self.commits = [] - def pending_layers(self): - return [(0, 0)] + def ready_layer(self): + return 0 def commit(self, *args): self.commits.append(args) @@ -2260,13 +2232,14 @@ class Transfer: def __init__(self, live, staging): self.live = live self.staging = staging - self.pending = [(0, 0)] + self.ready = 0 - def pending_layers(self): - return self.pending + def ready_layer(self): + return self.ready - def commit(self, layer, buffer_index, post_copy=None): - assert self.pending.pop(0) == (layer, buffer_index) + def commit(self, layer, post_copy=None): + assert self.ready == layer + self.ready = None self.live.copy_(self.staging, non_blocking=True) if post_copy is not None: post_copy() @@ -2707,6 +2680,33 @@ def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): assert calls == ["evaluation"] +def test_mode_backend_owns_eplb_poll_and_prefill_step(): + backend = object.__new__(ModeBackend) + calls = [] + backend.eplb_manager = SimpleNamespace( + poll=lambda: calls.append("poll"), + step=lambda: calls.append("step"), + ) + backend.prefill = lambda **_kwargs: calls.append("prefill") + + backend._poll_eplb() + backend._run_prefill(event_pack=object(), prefill_reqs=[]) + + assert calls == ["poll", "prefill", "step"] + + +def test_mode_backend_eplb_hooks_are_noops_when_disabled(): + backend = object.__new__(ModeBackend) + backend.eplb_manager = None + calls = [] + backend.prefill = lambda **_kwargs: calls.append("prefill") + + backend._poll_eplb() + backend._run_prefill(event_pack=object(), prefill_reqs=[]) + + assert calls == ["prefill"] + + def test_pinned_transfer_groups_sources_in_collective_order(): plan = [ TransferStep(0, 2, 1, 3), @@ -2780,9 +2780,8 @@ def record(self, _stream): transfer._release.set() transfer._consumed_event = Event() transfer._consumed_recorded = False - transfer._changed_dst_slots = () - transfer._pending = transfer_module.deque() - transfer._pending_lock = threading.Lock() + transfer._ready = None + transfer._ready_lock = threading.Lock() transfer._error = None transfer._thread = None copied = [] @@ -2792,19 +2791,19 @@ def record(self, _stream): transfer.start([(0, []), (1, [])]) deadline = time.monotonic() + 2 - while not transfer.pending_layers() and time.monotonic() < deadline: + while transfer.ready_layer() is None and time.monotonic() < deadline: time.sleep(0.001) - assert transfer.pending_layers() == [(0, 0)] + assert transfer.ready_layer() == 0 assert copied == [0] - transfer.commit(0, 0) + transfer.commit(0) deadline = time.monotonic() + 2 - while not transfer.pending_layers() and time.monotonic() < deadline: + while transfer.ready_layer() is None and time.monotonic() < deadline: time.sleep(0.001) - assert transfer.pending_layers() == [(1, 0)] + assert transfer.ready_layer() == 1 assert copied == [0, 1] assert transfer._consumed_event.synchronize_count == 1 - transfer.commit(1, 0) + transfer.commit(1) transfer.finish() @@ -2842,9 +2841,7 @@ def new_group(*args, **kwargs): monkeypatch.setattr( manager_module, "PinnedMemoryEPLBTransfer", - lambda weights, group, rank, world_size: ( - transfer_calls.append((weights, group, rank, world_size)) or transfer - ), + lambda weights, group, rank: (transfer_calls.append((weights, group, rank)) or transfer), ) logs = [] monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) @@ -2856,7 +2853,7 @@ def new_group(*args, **kwargs): manager.transfer_group, ) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 - assert transfer_calls == [([weight], groups[2], 0, 2)] + assert transfer_calls == [([weight], groups[2], 0)] assert manager.rebalance_gain_threshold == 0.07 assert "rebalance_gain_threshold=0.0700" in logs[0] assert manager._continuous_collection_start_step is None diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 5893cda007..372f1c10ce 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -55,11 +55,11 @@ def _free_port(): def _wait_for_ready_layer(transfer, control_group): deadline = time.monotonic() + 30 while time.monotonic() < deadline: - pending = transfer.pending_layers() - ready_count = torch.tensor([len(pending)], dtype=torch.int32) + ready_layer = transfer.ready_layer() + ready_count = torch.tensor([int(ready_layer is not None)], dtype=torch.int32) dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=control_group) if int(ready_count.item()) == 1: - return pending[0] + return ready_layer time.sleep(0.001) raise TimeoutError("EPLB transfer worker did not publish a globally ready layer") @@ -73,7 +73,7 @@ def _worker(rank, port): transfer_group = dist.new_group([0, 1], backend="gloo") weights = [_FakeWeight(rank, layer_index) for layer_index in range(2)] - transfer = PinnedMemoryEPLBTransfer(weights, transfer_group, rank, world_size=2) + transfer = PinnedMemoryEPLBTransfer(weights, transfer_group, rank) assert all(row.is_pinned() for _, row in transfer.pinned_rows) current = torch.tensor([[2], [0]]) target = torch.tensor([[3], [1]]) @@ -82,7 +82,7 @@ def _worker(rank, port): transfer.start(layer_plans) for expected_layer in range(2): - layer_index, buffer_index = _wait_for_ready_layer(transfer, control_group) + layer_index = _wait_for_ready_layer(transfer, control_group) assert layer_index == expected_layer expected_expert = 3 if rank == 0 else 1 expected_w13 = expected_layer * 100 + expected_expert @@ -91,7 +91,7 @@ def _worker(rank, port): assert torch.all(transfer.staging[1][1][0] == expected_w13 + 0.5) assert torch.all(transfer.staging[2][1][0] == expected_w2) assert torch.all(transfer.staging[3][1][0] == expected_w2 + 0.5) - transfer.commit(layer_index, buffer_index) + transfer.commit(layer_index) transfer.finish() torch.cuda.synchronize() From 0c8288f71270794299d3a70bdd9b2f4903844e4b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 03:09:22 +0000 Subject: [PATCH 22/72] refactor(eplb): extract extensible placement planner --- .../meta_weights/fused_moe/eplb_placement.py | 241 --- .../meta_weights/fused_moe/eplb_planner.py | 405 +++++ .../model_infer/mode_backend/eplb_manager.py | 584 +++---- .../model_infer/mode_backend/eplb_transfer.py | 33 +- unit_tests/common/fused_moe/test_eplb.py | 1439 ++--------------- 5 files changed, 778 insertions(+), 1924 deletions(-) create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py index c5ab5b0789..d043e4cd6f 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -1,7 +1,3 @@ -from typing import Dict, Tuple -import torch - - def build_initial_local_expert_ids( num_logical_experts: int, num_ranks: int, @@ -193,240 +189,3 @@ def build_logical_to_physical_maps_for_layers( ) for rank_to_logic_expert_ids in rank_to_logic_expert_ids_by_layer ] - - -def select_improving_placements( - expert_load: torch.Tensor, - current_placement: torch.Tensor, - candidate_placement: torch.Tensor, - *, - rebalance_gain_threshold: float, - expert_alignment: int | None = None, -) -> Tuple[torch.Tensor, torch.Tensor, Dict[str, float | int], torch.Tensor, torch.Tensor]: - """Select better layers and return current/final rank loads without re-estimation.""" - if not 0.0 <= rebalance_gain_threshold <= 1.0: - raise ValueError("rebalance_gain_threshold must be between 0.0 and 1.0") - assert current_placement.shape == candidate_placement.shape - current_rank_load = _estimate_rank_load(expert_load, current_placement, expert_alignment) - candidate_rank_load = _estimate_rank_load(expert_load, candidate_placement, expert_alignment) - current_critical = current_rank_load.max(dim=-1).values.sum(dim=0) - candidate_critical = candidate_rank_load.max(dim=-1).values.sum(dim=0) - # Each changed layer must reduce its own critical load. All selected - # changes must then collectively meet the configured model-level - # critical-load reduction threshold, avoiding low-gain migrations. - improved = candidate_critical < current_critical - selected = current_placement.clone() - selected[improved] = candidate_placement[improved] - selected_rank_load = torch.where(improved[:, None], candidate_rank_load, current_rank_load) - model_current_critical = current_critical.sum() - model_current_mean = current_rank_load.mean(dim=-1).sum() - model_selected_critical = selected_rank_load.max(dim=-1).values.sum() - model_selected_mean = selected_rank_load.mean(dim=-1).sum() - model_ratio = model_current_critical / model_current_mean.clamp_min(1.0) - candidate_model_ratio = model_selected_critical / model_selected_mean.clamp_min(1.0) - candidate_rebalance_gain = (model_current_critical - model_selected_critical) / model_current_critical.clamp_min( - 1.0 - ) - metrics = { - "model_imbalance_ratio": float(model_ratio.item()), - "candidate_model_imbalance_ratio": float(candidate_model_ratio.item()), - "candidate_rebalance_gain": float(candidate_rebalance_gain.item()), - "candidate_changed_layer_count": int(improved.sum().item()), - } - if candidate_rebalance_gain >= rebalance_gain_threshold: - return selected, improved, metrics, current_rank_load, selected_rank_load - return ( - current_placement.clone(), - torch.zeros_like(improved), - metrics, - current_rank_load, - current_rank_load, - ) - - -def plan_redundant_experts( - expert_load: torch.Tensor, - num_ranks: int, - num_redundant_experts_per_rank: int, - expert_alignment: int | None = None, - current_placement: torch.Tensor | None = None, - stickiness: float = 0.0, -) -> torch.Tensor: - """Plan replicas from [samples, layers, current_ranks, experts] loads. - - With ``current_placement`` and positive ``stickiness``, a candidate that - keeps an expert on its current rank receives a bonus of - ``stickiness * mean per-layer expert load``. This preserves rank - membership, not a particular redundant physical slot; target slots are - canonicalized against the current live rows before transfer and metadata - publication. A rank membership only changes when the move improves the - critical-load objective by more than that margin. - With zero stickiness, placement is determined solely by the load objective. - """ - if expert_alignment is not None: - assert expert_alignment > 0 - assert expert_load.ndim == 4 - _, num_layers, num_load_ranks, num_logical_experts = expert_load.shape - assert num_load_ranks in (1, num_ranks) - assert num_logical_experts % num_ranks == 0 - assert num_redundant_experts_per_rank > 0 - num_experts_per_rank = num_logical_experts // num_ranks - num_redundant = num_ranks * num_redundant_experts_per_rank - assert num_redundant <= num_logical_experts * (num_ranks - 1) - - load = expert_load.to(dtype=torch.float64, device="cpu") - placement = torch.full((num_layers, num_ranks, num_redundant_experts_per_rank), -1, dtype=torch.int64) - owner_rank = torch.arange(num_logical_experts, dtype=torch.int64) // num_experts_per_rank - if current_placement is not None: - assert tuple(current_placement.shape) == ( - num_layers, - num_ranks, - num_redundant_experts_per_rank, - ) - current_locations = _expert_locations(current_placement, num_logical_experts) - stickiness_scale = load.sum(dim=(0, 2, 3)) / num_logical_experts - else: - current_locations = None - stickiness_scale = None - - locations = _expert_locations(placement, num_logical_experts) - expert_rank = _expert_rank_load_all(load, locations, expert_alignment) - rank_load = expert_rank.sum(dim=2) - remaining_slots = torch.full((num_layers, num_ranks), num_redundant_experts_per_rank, dtype=torch.int64) - layer_indices = torch.arange(num_layers, dtype=torch.int64) - expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) - # Every iteration fills one slot per layer. Candidate expert evaluation - # is vectorized across all layers and logical experts, which keeps large - # GLM/Qwen planning comfortably on the CPU fast path. - for _ in range(num_redundant): - rank_order = torch.argsort(rank_load.sum(dim=0), dim=1, stable=True) - target_ranks = torch.full((num_layers,), -1, dtype=torch.int64) - legal = torch.zeros((num_layers, num_logical_experts), dtype=torch.bool) - for layer in range(num_layers): - for target_rank in rank_order[layer].tolist(): - if remaining_slots[layer, target_rank] == 0: - continue - candidate_legal = (owner_rank != target_rank) & ~locations[layer, :, target_rank] - if torch.any(candidate_legal): - target_ranks[layer] = target_rank - legal[layer] = candidate_legal - break - if torch.any(target_ranks < 0): - raise RuntimeError("EPLB planner found no valid redundant expert placement") - - candidate_locations = locations.clone() - candidate_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] = True - candidate_expert_rank = _expert_rank_load_all(load, candidate_locations, expert_alignment) - candidate_rank_load = rank_load[:, :, None, :] - expert_rank + candidate_expert_rank - critical = candidate_rank_load.max(dim=3).values.sum(dim=0) - critical.masked_fill_(~legal, torch.inf) - if current_locations is not None: - # An expert already held by the target rank is retained unless - # another candidate beats it by more than the stickiness margin. - # This is rank membership, not physical-slot stickiness. Masked - # (inf) candidates stay masked: inf - x == inf. - keep = current_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] - critical = critical - stickiness * stickiness_scale[:, None] * keep - selected_experts = critical.argmin(dim=1) - if torch.isinf(critical[layer_indices, selected_experts]).any(): - raise RuntimeError("EPLB planner found no valid redundant expert placement") - - slots = num_redundant_experts_per_rank - remaining_slots[layer_indices, target_ranks] - placement[layer_indices, target_ranks, slots] = selected_experts - selected_next = candidate_expert_rank[:, layer_indices, selected_experts] - selected_old = expert_rank[:, layer_indices, selected_experts] - rank_load += selected_next - selected_old - expert_rank[:, layer_indices, selected_experts] = selected_next - locations[layer_indices, selected_experts, target_ranks] = True - remaining_slots[layer_indices, target_ranks] -= 1 - - assert torch.all(placement >= 0) - return placement - - -def _estimate_rank_load( - expert_load: torch.Tensor, - rank_to_logic_expert_ids: torch.Tensor, - expert_alignment: int | None = None, -) -> torch.Tensor: - """Estimate [samples, layers, ranks] load from current-rank-local routing. - - Per-rank loads remain separate until assigned to physical replicas, then - combine before the per-expert alignment used by DeepEP. - """ - assert expert_load.ndim == 4 - _, num_layers, num_load_ranks, num_logical_experts = expert_load.shape - assert rank_to_logic_expert_ids.ndim == 3 and rank_to_logic_expert_ids.shape[0] == num_layers - num_ranks = rank_to_logic_expert_ids.shape[1] - assert num_load_ranks in (1, num_ranks) - assert num_logical_experts % num_ranks == 0 - if expert_alignment is not None: - assert expert_alignment > 0 - - rank_load = _expert_rank_load_all( - expert_load, - _expert_locations(rank_to_logic_expert_ids, num_logical_experts), - expert_alignment, - ).sum(dim=2) - return rank_load - - -def _expert_locations(rank_to_logic_expert_ids: torch.Tensor, num_logical_experts: int) -> torch.Tensor: - """Return ``[layer, logical expert, rank]`` physical-copy occupancy.""" - num_layers, num_ranks, num_redundant_experts_per_rank = rank_to_logic_expert_ids.shape - assert num_logical_experts % num_ranks == 0 - num_experts_per_rank = num_logical_experts // num_ranks - locations = torch.zeros( - (num_layers, num_logical_experts, num_ranks), - dtype=torch.bool, - device=rank_to_logic_expert_ids.device, - ) - expert_ids = torch.arange(num_logical_experts, device=locations.device) - owners = expert_ids // num_experts_per_rank - locations[:, expert_ids, owners] = True - layers = torch.arange(num_layers, device=locations.device)[:, None] - ranks = torch.arange(num_ranks, device=locations.device).repeat_interleave(num_redundant_experts_per_rank)[None, :] - flat_logic_expert_ids = rank_to_logic_expert_ids.reshape(num_layers, -1) - valid = flat_logic_expert_ids >= 0 - if torch.any(valid): - expanded_layers = layers.expand_as(flat_logic_expert_ids) - expanded_ranks = ranks.expand_as(flat_logic_expert_ids) - locations[ - expanded_layers[valid], - flat_logic_expert_ids[valid], - expanded_ranks[valid], - ] = True - return locations - - -def _current_rank_route(slots: torch.Tensor, num_load_ranks: int) -> torch.Tensor: - """当前 rank 有本地副本时只选本地,否则在所有副本之间均分。 - - ``num_load_ranks == 1`` 表示离线调用方只提供了聚合负载,此时无法判断 - 当前 rank,直接在全部副本之间均分。线上采样始终传入逐 rank 负载。 - """ - num_ranks = slots.shape[-1] - copies = slots.unsqueeze(-3).expand(*slots.shape[:-2], num_load_ranks, *slots.shape[-2:]) - if num_load_ranks == 1: - return copies.to(torch.float64) / copies.sum(dim=-1, keepdim=True) - - assert num_load_ranks == num_ranks - ranks = torch.arange(num_ranks, device=slots.device) - destination_rank_shape = (1,) * slots.ndim + (num_ranks,) - current_rank_shape = (1,) * (slots.ndim - 2) + (num_ranks, 1, 1) - local = copies & (ranks.reshape(destination_rank_shape) == ranks.reshape(current_rank_shape)) - selected = torch.where(local.any(dim=-1, keepdim=True), local, copies) - return selected.to(torch.float64) / selected.sum(dim=-1, keepdim=True) - - -def _expert_rank_load_all( - expert_load: torch.Tensor, - locations: torch.Tensor, - expert_alignment: int | None, -) -> torch.Tensor: - """Return aligned ``[samples, layers, expert, rank]`` contributions.""" - route = _current_rank_route(locations, expert_load.shape[2]) - physical_load = torch.einsum("slqe,lqer->sler", expert_load.to(torch.float64), route) - if expert_alignment is not None: - physical_load = torch.ceil(physical_load / expert_alignment) * expert_alignment - return physical_load diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py new file mode 100644 index 0000000000..84f5ac9625 --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py @@ -0,0 +1,405 @@ +"""Pure-Python redundant-expert placement planning for EPLB. + +The planner deliberately uses nested lists instead of tensors. Tensor +conversion belongs to the manager's distributed-communication and migration +boundaries; keeping it out of this module makes placement algorithms easy to +read, test, and replace. +""" + +from abc import ABC, abstractmethod +from collections import Counter +from dataclasses import dataclass +from math import ceil +from typing import Any, Dict, List + + +# [layer][logical expert] +LogicalExpertLoad = List[List[float]] +# [layer][rank][redundant slot] -> logical expert +ExpertPlacement = List[List[List[int]]] +# [layer][rank] +RankLoad = List[List[float]] + + +@dataclass +class EPLBPlan: + placement: ExpertPlacement + changed_layers: List[bool] + before_rank_load: RankLoad + after_rank_load: RankLoad + reason: str + minimum_layer_samples: float + required_layer_samples: float + + @property + def changed(self) -> bool: + return any(self.changed_layers) + + def as_dict(self) -> Dict[str, Any]: + before = _imbalance_summary(self.before_rank_load) + after = _imbalance_summary(self.after_rank_load) + before_critical = sum(max(layer) for layer in self.before_rank_load) + after_critical = sum(max(layer) for layer in self.after_rank_load) + gain = (before_critical - after_critical) / max(before_critical, 1.0) + return { + "kind": self.reason, + "placement": self.placement, + "changed_layers": self.changed_layers, + "minimum_layer_samples": self.minimum_layer_samples, + "required_layer_samples": self.required_layer_samples, + "before": before, + "after": after, + "rebalance_gain": gain, + "changed_layer_count": sum(self.changed_layers), + } + + +class EPLBPlanner(ABC): + """Interface for planning redundant expert placement.""" + + @abstractmethod + def plan( + self, + logical_expert_load: LogicalExpertLoad, + current_placement: ExpertPlacement, + ) -> EPLBPlan: + """Plan a concrete ``[layer][rank][redundant slot]`` placement.""" + + +class GreedyEPLBPlanner(EPLBPlanner): + """Greedy planner based on global logical-expert loads.""" + + def __init__( + self, + world_size: int, + num_redundant_experts_per_rank: int, + *, + expert_alignment: int = 1, + min_avg_tokens_per_expert: float = 0, + rebalance_gain_threshold: float = 0.0, + ): + if world_size <= 1: + raise ValueError("world_size must be greater than one") + if num_redundant_experts_per_rank <= 0: + raise ValueError("num_redundant_experts_per_rank must be positive") + if expert_alignment <= 0: + raise ValueError("expert_alignment must be positive") + if min_avg_tokens_per_expert < 0: + raise ValueError("min_avg_tokens_per_expert must be non-negative") + if not 0.0 <= rebalance_gain_threshold <= 1.0: + raise ValueError("rebalance_gain_threshold must be between 0.0 and 1.0") + self.world_size = world_size + self.num_redundant_experts_per_rank = num_redundant_experts_per_rank + self.expert_alignment = expert_alignment + self.min_avg_tokens_per_expert = min_avg_tokens_per_expert + self.rebalance_gain_threshold = rebalance_gain_threshold + + def plan( + self, + logical_expert_load: LogicalExpertLoad, + current_placement: ExpertPlacement, + ) -> EPLBPlan: + """Return a new concrete placement. + + ``logical_expert_load`` is ``[layer][logical_expert]`` and contains + the load summed across all ranks. + ``current_placement`` is ``[layer][rank][redundant_slot]``. The + returned placement has the same shape and already names the exact + physical slots that migration should update. + """ + load = [[float(value) for value in layer] for layer in logical_expert_load] + current = [[[int(expert) for expert in rank] for rank in layer] for layer in current_placement] + num_logical_experts = self._validate_inputs(load, current) + layer_samples = [sum(layer) for layer in load] + required_samples = self.min_avg_tokens_per_expert * num_logical_experts + minimum_samples = min(layer_samples) + before_rank_load = self.estimate_rank_load(load, current) + + if minimum_samples < required_samples: + return EPLBPlan( + placement=current, + changed_layers=[False] * len(load), + before_rank_load=before_rank_load, + after_rank_load=[layer[:] for layer in before_rank_load], + reason="insufficient", + minimum_layer_samples=minimum_samples, + required_layer_samples=required_samples, + ) + + candidates = [self._plan_layer(layer_load, current_layer) for layer_load, current_layer in zip(load, current)] + candidate_rank_load = self.estimate_rank_load(load, candidates) + changed_layers = [] + placement = [] + after_rank_load = [] + for layer, candidate in enumerate(candidates): + before = max(before_rank_load[layer]) + after = max(candidate_rank_load[layer]) + gain = (before - after) / max(before, 1.0) + changed = candidate != current[layer] and gain >= self.rebalance_gain_threshold + changed_layers.append(changed) + placement.append(candidate if changed else current[layer]) + after_rank_load.append(candidate_rank_load[layer][:] if changed else before_rank_load[layer][:]) + + return EPLBPlan( + placement=placement, + changed_layers=changed_layers, + before_rank_load=before_rank_load, + after_rank_load=after_rank_load, + reason="planned" if any(changed_layers) else "no_improvement", + minimum_layer_samples=minimum_samples, + required_layer_samples=required_samples, + ) + + def estimate_rank_load( + self, + logical_expert_load: LogicalExpertLoad, + placement: ExpertPlacement, + ) -> RankLoad: + """Estimate aligned physical-expert work for each layer and rank.""" + load = [[float(value) for value in layer] for layer in logical_expert_load] + normalized_placement = [[[int(expert) for expert in rank] for rank in layer] for layer in placement] + num_logical_experts = self._validate_inputs(load, normalized_placement) + estimated = [] + for layer_load, layer_placement in zip(load, normalized_placement): + locations = self._expert_locations(layer_placement, num_logical_experts) + rank_load = [0.0] * self.world_size + for expert, expert_locations in enumerate(locations): + for rank, value in enumerate(self._physical_load(layer_load[expert], expert_locations)): + rank_load[rank] += value + estimated.append(rank_load) + return estimated + + def _plan_layer( + self, + logical_load: List[float], + current_placement: List[List[int]], + ) -> List[List[int]]: + num_logical_experts = len(logical_load) + experts_per_rank = num_logical_experts // self.world_size + owner = [expert // experts_per_rank for expert in range(num_logical_experts)] + copy_count = self._allocate_copy_count(logical_load, owner) + + placement = [[-1] * self.num_redundant_experts_per_rank for _ in range(self.world_size)] + locations = [{owner_rank} for owner_rank in owner] + expert_rank_load = [ + self._physical_load( + logical_load[expert], + locations[expert], + ) + for expert in range(num_logical_experts) + ] + rank_load = [ + sum(expert_rank_load[expert][rank] for expert in range(num_logical_experts)) + for rank in range(self.world_size) + ] + instances = sorted( + [expert for expert, copies in enumerate(copy_count) for _ in range(copies - 1)], + key=lambda expert: ( + -copy_count[expert], + -logical_load[expert] / copy_count[expert], + expert, + ), + ) + + for index, expert in enumerate(instances): + best = None + for rank in range(self.world_size): + if rank in locations[expert]: + continue + empty_slots = [slot for slot, value in enumerate(placement[rank]) if value < 0] + if not empty_slots: + continue + slot = min( + empty_slots, + key=lambda candidate: ( + current_placement[rank][candidate] != expert, + candidate, + ), + ) + trial_placement = [row[:] for row in placement] + trial_placement[rank][slot] = expert + trial_locations = [set(ranks) for ranks in locations] + trial_locations[expert].add(rank) + if not self._can_complete( + instances[index + 1 :], + trial_placement, + trial_locations, + owner, + ): + continue + + next_expert_load = self._physical_load( + logical_load[expert], + trial_locations[expert], + ) + trial_rank_load = [ + value - expert_rank_load[expert][target_rank] + next_expert_load[target_rank] + for target_rank, value in enumerate(rank_load) + ] + candidate = ( + max(trial_rank_load), + current_placement[rank][slot] != expert, + sum(trial_rank_load), + rank, + slot, + trial_rank_load, + next_expert_load, + ) + if best is None or candidate[:5] < best[:5]: + best = candidate + + if best is None: + raise RuntimeError("EPLB planner found no valid redundant expert placement") + _, _, _, rank, slot, rank_load, next_expert_load = best + placement[rank][slot] = expert + locations[expert].add(rank) + expert_rank_load[expert] = next_expert_load + + return placement + + def _allocate_copy_count(self, logical_load: List[float], owner: List[int]) -> List[int]: + copy_count = [1] * len(logical_load) + replicas_by_owner = [0] * self.world_size + owner_capacity = self.num_redundant_experts_per_rank * (self.world_size - 1) + total_replicas = self.num_redundant_experts_per_rank * self.world_size + for _ in range(total_replicas): + candidates = [ + expert + for expert in range(len(logical_load)) + if copy_count[expert] < self.world_size and replicas_by_owner[owner[expert]] < owner_capacity + ] + if not candidates: + raise RuntimeError("EPLB planner cannot allocate all redundant copies") + expert = max( + candidates, + key=lambda candidate: ( + logical_load[candidate] / copy_count[candidate], + -candidate, + ), + ) + copy_count[expert] += 1 + replicas_by_owner[owner[expert]] += 1 + return copy_count + + def _can_complete( + self, + remaining_instances: List[int], + placement: List[List[int]], + locations: List[set], + owner: List[int], + ) -> bool: + """Check that a greedy choice leaves a legal assignment for all slots.""" + remaining = Counter(remaining_instances) + capacity = [sum(expert < 0 for expert in row) for row in placement] + memo = set() + + def search() -> bool: + if not remaining: + return True + state = ( + tuple(sorted(remaining.items())), + tuple(capacity), + tuple(tuple(sorted(ranks)) for ranks in locations), + ) + if state in memo: + return False + memo.add(state) + + expert = min( + remaining, + key=lambda item: ( + sum(capacity[rank] > 0 and rank not in locations[item] for rank in range(self.world_size)) + - remaining[item], + -remaining[item], + item, + ), + ) + candidate_ranks = [ + rank + for rank in range(self.world_size) + if capacity[rank] > 0 and rank not in locations[expert] and rank != owner[expert] + ] + if len(candidate_ranks) < remaining[expert]: + return False + count = remaining.pop(expert) + if count > 1: + remaining[expert] = count - 1 + for rank in sorted(candidate_ranks, key=lambda item: (-capacity[item], item)): + capacity[rank] -= 1 + locations[expert].add(rank) + if search(): + locations[expert].remove(rank) + capacity[rank] += 1 + remaining[expert] = count + return True + locations[expert].remove(rank) + capacity[rank] += 1 + remaining[expert] = count + return False + + return search() + + def _physical_load(self, logical_expert_load: float, locations: set) -> List[float]: + physical_expert_load = logical_expert_load / len(locations) + aligned_load = ceil(physical_expert_load / self.expert_alignment) * self.expert_alignment + result = [0.0] * self.world_size + for rank in locations: + result[rank] = aligned_load + return result + + def _expert_locations( + self, + placement: List[List[int]], + num_logical_experts: int, + ) -> List[set]: + experts_per_rank = num_logical_experts // self.world_size + locations = [{expert // experts_per_rank} for expert in range(num_logical_experts)] + for rank, row in enumerate(placement): + for expert in row: + if rank in locations[expert]: + raise ValueError(f"logical expert {expert} appears twice on rank {rank}") + locations[expert].add(rank) + return locations + + def _validate_inputs( + self, + logical_expert_load: LogicalExpertLoad, + placement: ExpertPlacement, + ) -> int: + if not logical_expert_load: + raise ValueError("logical_expert_load must contain at least one layer") + if len(placement) != len(logical_expert_load): + raise ValueError("load and placement must have the same number of layers") + num_logical_experts = len(logical_expert_load[0]) + if num_logical_experts == 0 or num_logical_experts % self.world_size: + raise ValueError("logical expert count must be positive and divisible by world_size") + if self.num_redundant_experts_per_rank > num_logical_experts - num_logical_experts // self.world_size: + raise ValueError("too many redundant slots to avoid local or duplicate replicas") + + for layer_load, layer_placement in zip(logical_expert_load, placement): + if len(layer_load) != num_logical_experts: + raise ValueError("each load layer must contain every logical expert") + if any(value < 0 for value in layer_load): + raise ValueError("logical expert load must be non-negative") + if len(layer_placement) != self.world_size or any( + len(rank) != self.num_redundant_experts_per_rank for rank in layer_placement + ): + raise ValueError("each placement layer must be [world_size][redundant_slots]") + if any(expert < 0 or expert >= num_logical_experts for rank in layer_placement for expert in rank): + raise ValueError("placement contains an invalid logical expert") + self._expert_locations(layer_placement, num_logical_experts) + return num_logical_experts + + +def _imbalance_summary(rank_load: RankLoad) -> Dict[str, float]: + values = sorted(value for layer in rank_load for value in layer) + if not values: + return {"max": 0.0, "p95": 0.0, "mean": 0.0, "ratio": 0.0} + mean = sum(values) / len(values) + p95 = values[ceil(0.95 * len(values)) - 1] + return { + "max": max(values), + "p95": p95, + "mean": mean, + "ratio": max(values) / max(mean, 1.0), + } diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 5e57df508e..1c5ba9deda 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -1,33 +1,32 @@ import threading import time -from typing import Dict, Optional import torch import torch.distributed as dist from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import ( - FusedMoeWeight, -) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( build_initial_local_expert_ids, build_logical_to_physical_maps_for_layers, - plan_redundant_experts, - select_improving_placements, ) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ( + EPLBPlanner, + GreedyEPLBPlanner, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import ( + FusedMoeWeight, +) +from lightllm.server.metrics.manager import MetricClient from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( PinnedMemoryEPLBTransfer, - align_target_placement, build_transfer_plan, ) -from lightllm.server.metrics.manager import MetricClient from lightllm.utils.dist_utils import ( get_global_rank, get_global_world_size, get_node_world_size, ) from lightllm.utils.envs_utils import ( - get_eplb_placement_stickiness, get_eplb_rebalance_gain_threshold, get_prefill_eplb_step_interval, ) @@ -38,12 +37,11 @@ EPLB_MIN_AVG_TOKENS_PER_EXPERT = 100 EPLB_EXPERT_ALIGNMENT = 128 EPLB_CONTROL_ERROR = -1 -EPLB_STEADY_SAMPLE_STEPS = 4 EPLB_EXPERT_IMBALANCE_RATIO_METRIC = "lightllm_eplb_topk_expert_imbalance_ratio" class EPLBManager: - """Online EPLB with asynchronous GPU expert migration.""" + """Collect logical load, invoke a planner, and publish migrated rows.""" def __init__(self, model: TpPartBaseModel): self.weights = _find_fused_moe_weights(model) @@ -52,100 +50,76 @@ def __init__(self, model: TpPartBaseModel): self.world_size = get_global_world_size() self.node_world_size = get_node_world_size() self._eplb_impls = [weight.fuse_moe_impl for weight in self.weights] - self.step_interval = get_prefill_eplb_step_interval() - self.rebalance_gain_threshold = get_eplb_rebalance_gain_threshold() - self.placement_stickiness = get_eplb_placement_stickiness() - self.sampling_interval = self.step_interval - self.prefill_steps = 0 - routed = {weight.fuse_moe_impl.n_routed_experts for weight in self.weights} + + routed = {impl.n_routed_experts for impl in self._eplb_impls} redundant = {impl.num_redundant_experts_per_rank for impl in self._eplb_impls} assert len(routed) == len(redundant) == 1 self.num_logical_experts = routed.pop() self.num_redundant_experts_per_rank = redundant.pop() - num_primary_experts_per_rank = self.num_logical_experts // self.world_size - initial_local_expert_ids_by_rank = build_initial_local_expert_ids( + self.step_interval = get_prefill_eplb_step_interval() + self.prefill_steps = 0 + self.next_evaluation_step = self.step_interval + + initial_local_expert_ids = build_initial_local_expert_ids( self.num_logical_experts, self.world_size, self.num_redundant_experts_per_rank, ) - initial_redundant_expert_ids_by_rank = [ - expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids_by_rank + num_primary_experts_per_rank = self.num_logical_experts // self.world_size + initial_redundant_expert_ids = [ + expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids ] self.current_placement = torch.tensor( - [initial_redundant_expert_ids_by_rank for _ in self.weights], + [initial_redundant_expert_ids for _ in self.weights], dtype=torch.int64, ) + self.planner: EPLBPlanner = GreedyEPLBPlanner( + self.world_size, + self.num_redundant_experts_per_rank, + expert_alignment=EPLB_EXPERT_ALIGNMENT, + min_avg_tokens_per_expert=EPLB_MIN_AVG_TOKENS_PER_EXPERT, + rebalance_gain_threshold=get_eplb_rebalance_gain_threshold(), + ) + self.in_flight = False + self.in_flight_layers = [] + self.in_flight_started_at = None self.target_placement = None self.target_metadata = None - self.in_flight_started_at = None self.evaluation_in_flight = False self._evaluation_lock = threading.Lock() self._evaluation_result = None self._evaluation_error = None self._evaluation_thread = None self.metric_client = None - # A fresh manager starts with one continuous base window. After a - # sufficient evaluation, steady state returns to the cheap sparse - # probe. An insufficient sparse probe schedules one fresh continuous - # base window before the next fixed sampling boundary. - self._continuous_collection_start_step: Optional[int] = None - self._continuous_collection_end_step: Optional[int] = self.step_interval - self._steady_collection_end_step: Optional[int] = None - self._reset_route_counters() - self._set_recording(True) - # Keep background evaluation collectives separate from the main-thread - # control/poll collectives: their ordering is intentionally independent. + self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") - # This control-group scalar is only touched from the main inference - # thread, never by the background evaluation thread. self._control_ready_count = torch.empty(1, dtype=torch.int32) self.transfer = PinnedMemoryEPLBTransfer(self.weights, self.transfer_group, self.global_rank) + self._restart_collection() + if self.global_rank == 0: logger.info( - "eplb enabled " - f"layers={len(self.weights)} num_logical_experts={self.num_logical_experts} " + f"eplb enabled layers={len(self.weights)} num_logical_experts={self.num_logical_experts} " f"num_redundant_experts_per_rank={self.num_redundant_experts_per_rank} " - f"step_interval={self.step_interval} " - f"rebalance_gain_threshold={self.rebalance_gain_threshold:.4f} " - f"placement_stickiness={self.placement_stickiness:.4f}" + f"step_interval={self.step_interval} planner={type(self.planner).__name__}" ) def poll(self): - """Poll only from a globally ordered pre-forward boundary.""" + """Advance EPLB only from a globally ordered pre-forward boundary.""" if self.in_flight: self._poll_in_flight() - return - if self.evaluation_in_flight and self._evaluation_ready_on_all_ranks(): - self._poll_evaluation() + elif self.evaluation_in_flight and self._evaluation_ready_on_all_ranks(): + self._finish_evaluation() def step(self): if self.in_flight or self.evaluation_in_flight: return self.prefill_steps += 1 - continuous_start = self._continuous_collection_start_step - continuous_end = self._continuous_collection_end_step - if continuous_end is not None: - if continuous_start is not None and self.prefill_steps == continuous_start: - self._set_recording(True) - if self.prefill_steps >= continuous_end: - self._start_evaluation() - return - sampling_interval = self.sampling_interval - phase = self.prefill_steps % sampling_interval - steady_collection_end_step = self._steady_collection_end_step - if steady_collection_end_step is not None: - if self.prefill_steps >= steady_collection_end_step: - self._steady_collection_end_step = None - self._start_evaluation() - return - if sampling_interval == 1: + if self.prefill_steps >= self.next_evaluation_step: self._start_evaluation() - return - if phase == sampling_interval - self._steady_sample_window_steps(): - self._arm_steady_collection(self.prefill_steps + self._steady_sample_window_steps()) def _set_recording(self, enabled: bool): for impl in self._eplb_impls: @@ -156,238 +130,99 @@ def _reset_route_counters(self): if counters: torch._foreach_zero_(counters) - def _control_count(self, value: int) -> torch.Tensor: - """Return the main-thread-only reusable control collective scalar.""" - return self._control_ready_count.fill_(value) - - def _clear_continuous_collection(self): - self._continuous_collection_start_step = None - self._continuous_collection_end_step = None - - def _steady_sample_window_steps(self) -> int: - return min(EPLB_STEADY_SAMPLE_STEPS, self.sampling_interval) - - def _arm_steady_collection(self, collection_end_step: int): - """Start the fixed sparse window without moving its evaluation boundary.""" + def _restart_collection(self): self._reset_route_counters() - self._steady_collection_end_step = collection_end_step self._set_recording(True) + self.next_evaluation_step = self.prefill_steps + self.step_interval - def _begin_continuous_collection(self): - minimum_end = self.prefill_steps + self.step_interval - collection_end = -(-minimum_end // self.sampling_interval) * self.sampling_interval - self._reset_route_counters() - self._steady_collection_end_step = None - self._continuous_collection_start_step = collection_end - self.step_interval - self._continuous_collection_end_step = collection_end - self._set_recording(self._continuous_collection_start_step == self.prefill_steps) - - def _prepare_next_sampling_window(self): - """Clear the current window and arm the next sparse sampling window.""" - self._clear_continuous_collection() - self._steady_collection_end_step = None - if self.sampling_interval == 1: - self._reset_route_counters() - self._set_recording(True) - elif self.sampling_interval <= EPLB_STEADY_SAMPLE_STEPS: - # There is no later pre-boundary manager step at which to arm a - # full clamped window, so arm immediately but keep the same next - # fixed boundary. - self._arm_steady_collection(self.prefill_steps + self.sampling_interval) - else: - self._reset_route_counters() - self._set_recording(False) - - def _collect_local_samples(self) -> torch.Tensor: + def _collect_local_load(self) -> torch.Tensor: counters = [impl.route_counter for impl in self._eplb_impls] if any(counter.ndim != 1 or counter.shape[0] != self.num_logical_experts for counter in counters): raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") - return torch.stack(counters).unsqueeze(0).cpu() - - def _publish_expert_load_metrics(self, result): - if self.global_rank != 0 or "expert_imbalance_ratio" not in result: - return - if self.metric_client is None: - self.metric_client = MetricClient(get_shm_port_args().metric_port) - self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) - - def _commit_layer_metadata(self, layer_index: int): - impl = self._eplb_impls[layer_index] - impl.logical_to_physical_map.copy_(self.target_metadata[layer_index], non_blocking=True) + return torch.stack(counters).cpu() - def _finish_rebalance(self): - self.current_placement = self.target_placement - self.target_placement = None - self.target_metadata = None - self.in_flight = False - self._prepare_next_sampling_window() - if self.global_rank == 0: - logger.info(f"eplb completed wall_time={time.time() - self.in_flight_started_at:.2f}s") - - def _poll_in_flight(self): - local_error = None - try: - ready_layer = self.transfer.ready_layer() - except BaseException as exc: - ready_layer = None - local_error = exc - ready_count = self._control_count( - EPLB_CONTROL_ERROR if local_error is not None else int(ready_layer is not None) - ) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) - ready_count = int(ready_count.item()) - if ready_count < 0: - if local_error is not None: - raise RuntimeError("EPLB transfer worker failed on this rank") from local_error - raise RuntimeError("EPLB transfer worker failed on another rank") - if ready_count == 0: - return - if ready_layer != self.in_flight_layers[0]: - raise RuntimeError(f"EPLB ready layer {ready_layer} does not match expected {self.in_flight_layers[0]}") - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - # Previous forward is queued on the shared overlap stream; order the - # live-weight commit after it. The subsequent wait orders the next forward. - torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - self.transfer.commit( - ready_layer, - lambda: self._commit_layer_metadata(ready_layer), - ) - self.in_flight_layers.pop(0) - if not self.in_flight_layers: - self.transfer.finish() - self._finish_rebalance() + def _control_count(self, value: int) -> torch.Tensor: + return self._control_ready_count.fill_(value) def _plan_and_broadcast(self, global_load: torch.Tensor): - """Plan on rank zero and share the serializable result on the evaluation group.""" result = None local_error = None if self.global_rank == 0: try: - minimum = self.num_logical_experts * EPLB_MIN_AVG_TOKENS_PER_EXPERT - layer_samples = global_load.sum(dim=(0, 2, 3)) - if torch.any(layer_samples < minimum): - result = { - "kind": "insufficient", - "minimum_layer_samples": int(layer_samples.min().item()), - "minimum": minimum, - } - else: - candidate = plan_redundant_experts( - global_load, - self.world_size, - self.num_redundant_experts_per_rank, - expert_alignment=EPLB_EXPERT_ALIGNMENT, - current_placement=self.current_placement, - stickiness=self.placement_stickiness, - ) - placement, improved, metrics, before_load, after_load = select_improving_placements( - global_load, - self.current_placement, - candidate, - expert_alignment=EPLB_EXPERT_ALIGNMENT, - rebalance_gain_threshold=self.rebalance_gain_threshold, - ) - if bool(torch.any(improved)): - # A planner placement identifies experts by rank, not - # by redundant slot. Canonicalize every selected row - # before broadcasting so transfer, metadata, and the - # next current_placement all describe the same live - # physical expert rows. - placement = placement.clone() - for layer_index in torch.nonzero(improved, as_tuple=False).flatten().tolist(): - placement[layer_index] = align_target_placement( - self.current_placement[layer_index], - placement[layer_index], - ) - result = { - "kind": ("planned" if bool(torch.any(improved)) else "no_improvement"), - "placement": placement, - "improved": improved, - "before": _imbalance_summary(before_load), - "after": _imbalance_summary(after_load), - **metrics, - } + result = self.planner.plan( + global_load.tolist(), + self.current_placement.tolist(), + ).as_dict() except BaseException as exc: local_error = exc result = {"kind": "error", "message": f"{type(exc).__name__}: {exc}"} if self.world_size > 1: - result_list = [result] - dist.broadcast_object_list(result_list, src=0, group=self.evaluation_group) - result = result_list[0] + values = [result] + dist.broadcast_object_list(values, src=0, group=self.evaluation_group) + result = values[0] if result["kind"] == "error": if local_error is not None: raise RuntimeError("EPLB planner failed on rank zero") from local_error raise RuntimeError(f"EPLB planner failed on rank zero: {result['message']}") return result + def _build_rebalance_data(self, result): + metadata = [None] * len(self.weights) + layer_plans = [] + changed_layer_indices = [layer_index for layer_index, changed in enumerate(result["changed_layers"]) if changed] + if not changed_layer_indices: + return metadata, layer_plans + + num_primary_experts_per_rank = self.num_logical_experts // self.world_size + rank_to_logic_expert_ids_by_layer = [] + for layer_index in changed_layer_indices: + rank_to_logic_expert_ids_by_layer.append( + [ + list( + range( + rank * num_primary_experts_per_rank, + (rank + 1) * num_primary_experts_per_rank, + ) + ) + + result["placement"][layer_index][rank] + for rank in range(self.world_size) + ] + ) + maps = torch.tensor( + build_logical_to_physical_maps_for_layers( + rank_to_logic_expert_ids_by_layer, + self.num_logical_experts, + current_rank=self.global_rank, + ), + dtype=torch.int32, + ) + for offset, layer_index in enumerate(changed_layer_indices): + target = torch.tensor(result["placement"][layer_index], dtype=torch.int64) + metadata[layer_index] = maps[offset] + layer_plans.append( + ( + layer_index, + build_transfer_plan( + self.current_placement[layer_index], + target, + self.num_logical_experts, + self.world_size, + self.node_world_size, + ), + ) + ) + return metadata, layer_plans + def _evaluate_after_event(self, event: torch.cuda.Event): - """Run the CPU/Gloo planning phase after the frozen CUDA counters are ready.""" try: torch.cuda.set_device(self._eplb_impls[0].route_counter.device) event.synchronize() - local_load = self._collect_local_samples() - sample_window_steps = ( - self.step_interval - if self._continuous_collection_end_step is not None - else self._steady_sample_window_steps() - ) - # 保留每个当前 rank 的负载,planner 才能准确模拟“本卡优先, - # 否则在所有远端副本间分配”的运行时路由规则。 - global_load = torch.zeros( - (*local_load.shape[:2], self.world_size, local_load.shape[2]), - dtype=local_load.dtype, - ) - global_load[:, :, self.global_rank] = local_load + global_load = self._collect_local_load() dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) result = self._plan_and_broadcast(global_load) result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) - result["sample_window_steps"] = sample_window_steps if result["kind"] == "planned": - metadata = [None] * len(self.weights) - layer_plans = [] - improved_layer_indices = torch.nonzero(result["improved"], as_tuple=False).flatten() - if improved_layer_indices.numel(): - redundant_placements = result["placement"][improved_layer_indices].tolist() - num_primary_experts_per_rank = self.num_logical_experts // self.world_size - rank_to_logic_expert_ids_by_layer = [ - [ - list( - range( - rank * num_primary_experts_per_rank, - (rank + 1) * num_primary_experts_per_rank, - ) - ) - + rank_redundant_expert_ids - for rank, rank_redundant_expert_ids in enumerate(layer_placement) - ] - for layer_placement in redundant_placements - ] - maps_for_improved_layers = torch.tensor( - build_logical_to_physical_maps_for_layers( - rank_to_logic_expert_ids_by_layer, - self.num_logical_experts, - current_rank=self.global_rank, - ), - dtype=torch.int32, - ) - for improved_layer_offset, layer_index in enumerate(improved_layer_indices.tolist()): - placement = result["placement"][layer_index] - metadata[layer_index] = maps_for_improved_layers[improved_layer_offset] - layer_plans.append( - ( - layer_index, - build_transfer_plan( - self.current_placement[layer_index], - placement, - self.num_logical_experts, - self.world_size, - self.node_world_size, - ), - ) - ) - result["metadata"] = metadata - result["layer_plans"] = layer_plans + result["metadata"], result["layer_plans"] = self._build_rebalance_data(result) with self._evaluation_lock: self._evaluation_result = result except BaseException as exc: @@ -405,159 +240,134 @@ def _start_evaluation(self): self._evaluation_thread = threading.Thread(target=self._evaluate_after_event, args=(event,), daemon=True) self._evaluation_thread.start() - def _poll_evaluation(self): - if not self.evaluation_in_flight: - return False + def _evaluation_ready_on_all_ranks(self) -> bool: with self._evaluation_lock: error = self._evaluation_error result = self._evaluation_result - if error is not None or result is not None: - self._evaluation_result = None - self._evaluation_error = None - if error is not None: - self._evaluation_thread.join() - self.evaluation_in_flight = False - self._evaluation_thread = None - raise error - if result is None: - return True + status = EPLB_CONTROL_ERROR if error is not None else int(result is not None) + ready_count = self._control_count(status) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) + ready_count = int(ready_count.item()) + if ready_count < 0: + if error is not None: + raise RuntimeError("EPLB evaluation failed on this rank") from error + raise RuntimeError("EPLB evaluation failed on another rank") + return bool(ready_count) + + def _finish_evaluation(self): + with self._evaluation_lock: + error = self._evaluation_error + result = self._evaluation_result + self._evaluation_error = None + self._evaluation_result = None self._evaluation_thread.join() - self.evaluation_in_flight = False self._evaluation_thread = None - self._publish_expert_load_metrics(result) - if result["kind"] == "insufficient": - from_continuous_window = self._continuous_collection_end_step is not None - if from_continuous_window: - self.sampling_interval = min(self.sampling_interval * 4, self.step_interval * 16) - self._prepare_next_sampling_window() - else: - self._begin_continuous_collection() - if self.global_rank == 0: - if from_continuous_window: - logger.info( - "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " - "next_sampling_interval=%s sample_window_steps=%s", - self.prefill_steps, - result["minimum_layer_samples"], - result["minimum"], - self.sampling_interval, - result.get("sample_window_steps"), - ) - else: - logger.info( - "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " - "scheduled_fresh_window_start=%s scheduled_fresh_window_end=%s " - "sample_window_steps=%s", - self.prefill_steps, - result["minimum_layer_samples"], - result["minimum"], - self._continuous_collection_start_step, - self._continuous_collection_end_step, - result.get("sample_window_steps"), - ) - return False - if result["kind"] == "no_improvement": - self.sampling_interval = min(self.sampling_interval * 4, self.step_interval * 16) + self.evaluation_in_flight = False + if error is not None: + raise error + self._publish_expert_load_metric(result) + if result["kind"] != "planned": if self.global_rank == 0: logger.info( - "eplb skip rearrangement: no model improvement model_imbalance_ratio=%.4f " - "candidate_model_imbalance_ratio=%.4f candidate_rebalance_gain=%.4f " - "candidate_changed_layer_count=%s actual_changed_layer_count=0 next_sampling_interval=%s " - "sample_window_steps=%s", - result["model_imbalance_ratio"], - result["candidate_model_imbalance_ratio"], - result["candidate_rebalance_gain"], - result["candidate_changed_layer_count"], - self.sampling_interval, - result.get("sample_window_steps"), + "eplb skip rearrangement kind=%s minimum_layer_samples=%s required_layer_samples=%s", + result["kind"], + result["minimum_layer_samples"], + result["required_layer_samples"], ) - self._prepare_next_sampling_window() - return False + self._restart_collection() + return self._start_rebalance(result) - return True - def _evaluation_ready_on_all_ranks(self) -> bool: - with self._evaluation_lock: - local_error = self._evaluation_error - local_result = self._evaluation_result - local_status = EPLB_CONTROL_ERROR if local_error is not None else int(local_result is not None) - ready_count = self._control_count(local_status) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) - ready_count = int(ready_count.item()) - if ready_count < 0: - if local_error is not None: - raise RuntimeError("EPLB evaluation failed on this rank") from local_error - raise RuntimeError("EPLB evaluation failed on another rank") - return bool(ready_count) + def _publish_expert_load_metric(self, result): + if self.global_rank != 0: + return + if self.metric_client is None: + self.metric_client = MetricClient(get_shm_port_args().metric_port) + self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) def _start_rebalance(self, result): - placement = result["placement"] - layer_plans = result["layer_plans"] - self.sampling_interval = self.step_interval - self._clear_continuous_collection() - self._reset_route_counters() - self.target_placement = placement + self.target_placement = torch.tensor(result["placement"], dtype=torch.int64) self.target_metadata = result["metadata"] - self.in_flight_layers = [layer_index for layer_index, _ in layer_plans] + self.in_flight_layers = [layer_index for layer_index, _ in result["layer_plans"]] self.in_flight = True self.in_flight_started_at = time.time() - self.transfer.start(layer_plans) + self.transfer.start(result["layer_plans"]) if self.global_rank == 0: - actual_changed_slot_count = sum(len(plan) for _, plan in layer_plans) - cross_node_transfer_count = sum( - step.src_rank // self.node_world_size != step.dst_rank // self.node_world_size - for _, plan in layer_plans - for step in plan - ) + changed_slots = sum(len(plan) for _, plan in result["layer_plans"]) logger.info( - "eplb started prefill_steps=%s max_before=%.4f max_after=%.4f p95_before=%.4f p95_after=%.4f " - "model_imbalance_ratio=%.4f candidate_model_imbalance_ratio=%.4f " - "candidate_rebalance_gain=%.4f candidate_changed_layer_count=%s " - "actual_changed_layer_count=%s actual_changed_slot_count=%s cross_node_transfer_count=%s " - "sample_window_steps=%s", + "eplb started prefill_steps=%s max_before=%.4f max_after=%.4f " + "p95_before=%.4f p95_after=%.4f rebalance_gain=%.4f " + "changed_layer_count=%s changed_slot_count=%s", self.prefill_steps, result["before"]["max"], result["after"]["max"], result["before"]["p95"], result["after"]["p95"], - result["model_imbalance_ratio"], - result["candidate_model_imbalance_ratio"], - result["candidate_rebalance_gain"], - result["candidate_changed_layer_count"], - len(layer_plans), - actual_changed_slot_count, - cross_node_transfer_count, - result.get("sample_window_steps"), + result["rebalance_gain"], + result["changed_layer_count"], + changed_slots, ) + def _poll_in_flight(self): + local_error = None + try: + ready_layer = self.transfer.ready_layer() + except BaseException as exc: + ready_layer = None + local_error = exc + ready_count = self._control_count( + EPLB_CONTROL_ERROR if local_error is not None else int(ready_layer is not None) + ) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) + ready_count = int(ready_count.item()) + if ready_count < 0: + if local_error is not None: + raise RuntimeError("EPLB transfer worker failed on this rank") from local_error + raise RuntimeError("EPLB transfer worker failed on another rank") + if ready_count == 0: + return + if ready_layer != self.in_flight_layers[0]: + raise RuntimeError(f"EPLB ready layer {ready_layer} does not match expected {self.in_flight_layers[0]}") + + from lightllm.server.router.model_infer.infer_batch import g_infer_context -def _imbalance_summary(rank_load: torch.Tensor) -> Dict[str, float]: - if rank_load.ndim != 3: - raise ValueError("rank_load must be [samples, layers, ranks]") - critical = rank_load.max(dim=2).values.sum(dim=0) - mean = rank_load.mean(dim=2).sum(dim=0) - layer_imbalance = critical / mean.clamp_min(1.0) - sorted_imbalance = torch.sort(layer_imbalance).values - p95_index = max(0, (95 * layer_imbalance.numel() + 99) // 100 - 1) - return { - "max": float(layer_imbalance.max().item()), - "p95": float(sorted_imbalance[p95_index].item()), - } + torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) + self.transfer.commit(ready_layer, lambda: self._commit_layer_metadata(ready_layer)) + self.in_flight_layers.pop(0) + if not self.in_flight_layers: + self.transfer.finish() + self._finish_rebalance() + + def _commit_layer_metadata(self, layer_index: int): + self._eplb_impls[layer_index].logical_to_physical_map.copy_( + self.target_metadata[layer_index], + non_blocking=True, + ) + + def _finish_rebalance(self): + self.current_placement = self.target_placement + self.target_placement = None + self.target_metadata = None + self.in_flight = False + self._restart_collection() + if self.global_rank == 0: + logger.info( + "eplb completed wall_time=%.2fs", + time.time() - self.in_flight_started_at, + ) def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: """Average each layer's maximum-to-mean logical-expert token ratio.""" - if global_load.ndim != 4: - raise ValueError("global_load must be [samples, layers, ranks, logical_experts]") - if global_load.shape[1] == 0 or global_load.shape[3] == 0: - raise ValueError("global_load must contain at least one layer and logical expert") - layer_expert_load = global_load.sum(dim=(0, 2)).to(torch.float64) - layer_means = layer_expert_load.mean(dim=1) + if global_load.ndim != 2: + raise ValueError("global_load must be [layers, logical_experts]") + global_load = global_load.to(torch.float64) + layer_means = global_load.mean(dim=1) valid_layers = layer_means > 0 if not torch.any(valid_layers): return 0.0 - layer_ratios = layer_expert_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] - return float(layer_ratios.mean().item()) + ratios = global_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] + return float(ratios.mean().item()) def _find_fused_moe_weights(model): diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index b6c4f69c74..32477f1687 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -1,7 +1,7 @@ """Layer-by-layer expert-row migration for EPLB.""" import threading -from collections import Counter, defaultdict +from collections import defaultdict from dataclasses import dataclass from typing import List, Sequence, Tuple @@ -19,33 +19,6 @@ class TransferStep: src_local_row: int -def align_target_placement(current: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - """Canonicalize a target row layout without moving retained experts.""" - assert current.ndim == target.ndim == 2 - assert tuple(current.shape) == tuple(target.shape) - - aligned_target_rows = [] - for current_row, target_row in zip(current.tolist(), target.tolist()): - remaining_target = Counter(target_row) - aligned_row = list(current_row) - freed_slots = [] - for slot, expert in enumerate(current_row): - if remaining_target[expert] > 0: - remaining_target[expert] -= 1 - else: - freed_slots.append(slot) - new_experts = [] - for expert in target_row: - if remaining_target[expert] > 0: - new_experts.append(expert) - remaining_target[expert] -= 1 - assert len(freed_slots) == len(new_experts) - for slot, expert in zip(freed_slots, new_experts): - aligned_row[slot] = expert - aligned_target_rows.append(aligned_row) - return target.new_tensor(aligned_target_rows) - - def build_transfer_plan( current: torch.Tensor, target: torch.Tensor, @@ -56,7 +29,7 @@ def build_transfer_plan( assert tuple(current.shape) == tuple(target.shape) == (world_size, current.shape[1]) num_experts_per_rank = num_logical_experts // world_size current_rows = current.tolist() - aligned_target_rows = align_target_placement(current, target).tolist() + target_rows = target.tolist() # A primary row is always a valid source. Existing replicas are also # candidates so that a destination can prefer a same-node copy. candidates_by_expert = [ @@ -68,7 +41,7 @@ def build_transfer_plan( source_load = [0] * world_size plan = [] for dst_rank in range(world_size): - for dst_slot, expert in enumerate(aligned_target_rows[dst_rank]): + for dst_slot, expert in enumerate(target_rows[dst_rank]): if expert == current_rows[dst_rank][dst_slot]: continue src_rank, src_row = min( diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index c107590e2a..853d46f466 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -10,9 +10,10 @@ build_initial_local_expert_ids, build_logical_to_physical_map, build_logical_to_physical_maps_for_layers, - _estimate_rank_load, - plan_redundant_experts, - select_improving_placements, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ( + EPLBPlanner, + GreedyEPLBPlanner, ) from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs @@ -43,7 +44,6 @@ PinnedMemoryEPLBTransfer, TransferStep, _commit_staging_rows, - align_target_placement, build_transfer_plan, ) @@ -114,34 +114,6 @@ def _rank_to_logic_expert_ids(redundant_placement, num_logical_experts): ] -def _manual_runtime_rank_load(source_load, placement, alignment): - """Reference the committed runtime logical-to-physical maps on CPU.""" - samples, layers, source_ranks, num_logical_experts = source_load.shape - ranks, redundant = placement.shape[1:] - assert source_ranks == ranks - num_experts_per_rank = num_logical_experts // ranks - num_physical_experts_per_rank = num_experts_per_rank + redundant - raw = torch.zeros((samples, layers, ranks, num_logical_experts), dtype=torch.float64) - for layer in range(layers): - for current_rank in range(ranks): - logical_to_physical = build_logical_to_physical_map( - _rank_to_logic_expert_ids(placement[layer].tolist(), num_logical_experts), - num_logical_experts, - current_rank=current_rank, - ) - for expert, packed_row in enumerate(logical_to_physical): - num_replicas = packed_row[0] - physical_expert_ids = packed_row[2 : num_replicas + 2] - if packed_row[1]: - physical_expert_ids = physical_expert_ids[:1] - for physical_id in physical_expert_ids: - rank = physical_id // num_physical_experts_per_rank - raw[:, layer, rank, expert] += source_load[:, layer, current_rank, expert] / len( - physical_expert_ids - ) - return (torch.ceil(raw / alignment) * alignment).sum(dim=3) - - def test_base_call_template_forwards_selection_and_capture_callback(): class Impl(FuseMoeBaseImpl): def _select_experts( @@ -310,307 +282,129 @@ def test_build_initial_local_expert_ids_rejects_local_or_duplicate_replicas(): build_initial_local_expert_ids(8, 4, 7) -def test_fused_moe_loads_default_replicas_into_their_physical_rows(): - weight = object.__new__(fused_weight_module.FusedMoeWeight) - weight.lock = threading.Lock() - loaded = [] +def test_eplb_planner_defines_an_abstract_planning_interface(): + with pytest.raises(TypeError): + EPLBPlanner() - def load_weight(expert, local, _weights): - loaded.append(("weight", expert, local)) + assert isinstance(GreedyEPLBPlanner(2, 1), EPLBPlanner) - def load_scale(expert, local, _weights): - loaded.append(("scale", expert, local)) - - def load_zero_point(expert, local, _weights): - loaded.append(("zero", expert, local)) - weight._load_expert = load_weight - weight._load_expert_scale = load_scale - weight._load_expert_zero_point = load_zero_point - local_logic_expert_ids_list = build_initial_local_expert_ids(8, 4, 2)[0] - - weight._load_weight(local_logic_expert_ids_list, {}) - - assert local_logic_expert_ids_list == [0, 1, 2, 3] - assert [entry for entry in loaded if entry[0] == "weight"] == [ - ("weight", 0, 0), - ("weight", 1, 1), - ("weight", 2, 2), - ("weight", 3, 3), - ] - - -def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): - expert_load = ( - torch.tensor( - [ - [100, 90, 80, 70, 60, 50, 40, 30], - [30, 40, 50, 60, 70, 80, 90, 100], - ] - ) - .unsqueeze(0) - .unsqueeze(2) - ) - placement = plan_redundant_experts(expert_load, num_ranks=4, num_redundant_experts_per_rank=2) - - for layer_placement in placement: - for rank, expert_ids in enumerate(layer_placement.tolist()): - assert len(expert_ids) == len(set(expert_ids)) - assert all(expert_id // 2 != rank for expert_id in expert_ids) - - -def test_plan_redundant_experts_minimizes_samplewise_aligned_critical_load(): - samples = torch.tensor([[[300, 20, 20, 200]], [[100, 300, 40, 160]]]).unsqueeze(2) - placement = plan_redundant_experts(samples, num_ranks=2, num_redundant_experts_per_rank=1, expert_alignment=128) - candidates = [torch.tensor([[[left], [right]]]) for left in (2, 3) for right in (0, 1)] - - def critical(candidate): - return _estimate_rank_load(samples, candidate, expert_alignment=128).max(dim=2).values.sum() - - assert torch.equal(placement, torch.tensor([[[3], [0]]])) - assert critical(placement) == min(critical(candidate) for candidate in candidates) - - -def test_select_improving_placements_rejects_regressing_layer(): - expert_load = torch.tensor([[8649, 5740, 5002, 3441]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - regressing_candidate = torch.tensor([[[1], [0]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, regressing_candidate, rebalance_gain_threshold=0.05 - ) - - current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() - candidate_ratio = ( - _estimate_rank_load(expert_load, regressing_candidate).max() - / _estimate_rank_load(expert_load, regressing_candidate).mean() +def test_eplb_planner_builds_legal_concrete_slot_layout(): + planner = GreedyEPLBPlanner( + 4, + 1, + expert_alignment=1, + min_avg_tokens_per_expert=0, + rebalance_gain_threshold=0.0, ) - assert current_ratio.item() == pytest.approx(1.1007, abs=1e-4) - assert candidate_ratio.item() == pytest.approx(1.1184, abs=1e-4) - assert not improved.item() - assert torch.equal(selected, current) + current = _initial_extra_expert_placement(8, 4, 1).unsqueeze(0).tolist() + load = torch.ones((1, 4, 8), dtype=torch.int64) + load[:, :, 0] = 1000 + load[:, :, 4] = 500 + result = planner.plan(load.sum(dim=1).tolist(), current) + placement = result.placement[0] -def test_select_improving_placements_rejects_near_balance_when_gain_is_below_threshold(): - expert_load = torch.tensor([[1, 2, 1, 17]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[3], [0]]]) - candidate = torch.tensor([[[3], [1]]]) + for rank, row in enumerate(placement): + assert len(row) == len(set(row)) + assert all(expert // 2 != rank for expert in row) + assert max(map(max, result.after_rank_load)) <= max(map(max, result.before_rank_load)) + assert isinstance(result.placement, list) + assert isinstance(result.changed_layers, list) - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) - current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() - candidate_ratio = ( - _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() - ) - assert current_ratio.item() == pytest.approx(1.047619, abs=1e-6) - assert candidate_ratio.item() == pytest.approx(1.0) - assert not improved.item() - assert torch.equal(selected, current) - - -def test_select_improving_placements_accepts_alignment_aware_gain_even_when_current_ranks_are_balanced(): - expert_load = torch.tensor([[100, 129, 100, 129]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[3], [1]]]) - - selected, improved, metrics, _before_load, _after_load = select_improving_placements( - expert_load, - current, - candidate, - rebalance_gain_threshold=0.05, +def test_eplb_planner_estimator_distributes_global_load_across_copies(): + planner = GreedyEPLBPlanner( + 4, + 1, expert_alignment=128, + min_avg_tokens_per_expert=0, ) + placement = [[[2], [4], [6], [0]]] + load = [[100, 200, 300, 400, 500, 600, 700, 800]] - assert metrics["model_imbalance_ratio"] == pytest.approx(1.0) - assert metrics["candidate_rebalance_gain"] == pytest.approx(0.25) - assert improved.item() - assert torch.equal(selected, candidate) - - -def test_select_improving_placements_rejects_insufficient_rebalance_gain(): - expert_load = torch.tensor([[1, 1, 6, 7]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[3], [0]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) - - current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() - candidate_ratio = ( - _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() - ) - relative_improvement = (current_ratio - candidate_ratio) / current_ratio - assert current_ratio.item() == pytest.approx(1.4) - assert candidate_ratio.item() == pytest.approx(1.333333, abs=1e-6) - assert relative_improvement.item() == pytest.approx(0.047619, abs=1e-6) - assert not improved.item() - assert torch.equal(selected, current) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, - current, - candidate, - rebalance_gain_threshold=0.04, - ) - - assert improved.item() - assert torch.equal(selected, candidate) - - -@pytest.mark.parametrize("rebalance_gain_threshold", [-0.01, 1.01, float("nan"), float("inf")]) -def test_select_improving_placements_rejects_invalid_rebalance_gain_threshold( - rebalance_gain_threshold, -): - expert_load = torch.tensor([[1, 1, 1, 2]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[3], [0]]]) - - with pytest.raises(ValueError, match="rebalance_gain_threshold"): - select_improving_placements( - expert_load, - current, - candidate, - rebalance_gain_threshold=rebalance_gain_threshold, - ) + predicted = planner.estimate_rank_load(load, placement) + assert predicted == [[640, 1024, 1280, 1408]] -def test_select_improving_placements_accepts_sufficient_rebalance_gain(): - expert_load = torch.tensor([[1, 1, 1, 2]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[3], [0]]]) - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) +def test_eplb_planner_rejects_an_under_sampled_window(): + planner = GreedyEPLBPlanner(2, 1, min_avg_tokens_per_expert=100) + current = [[[2], [0]]] + load = [[1, 1, 1, 1]] - current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() - candidate_ratio = ( - _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() - ) - relative_improvement = (current_ratio - candidate_ratio) / current_ratio - assert current_ratio.item() == pytest.approx(1.2) - assert candidate_ratio.item() == pytest.approx(1.0) - assert relative_improvement.item() == pytest.approx(1 / 6) - assert improved.item() - assert torch.equal(selected, candidate) + result = planner.plan(load, current) + assert result.reason == "insufficient" + assert not result.changed + assert result.placement == current -def test_select_improving_placements_rejects_raw_improvement_that_does_not_improve_aligned_compute(): - expert_load = torch.tensor([[1, 1, 1, 8]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - raw_improving_candidate = torch.tensor([[[3], [0]]]) - _, raw_improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, raw_improving_candidate, rebalance_gain_threshold=0.05 +def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): + planner = GreedyEPLBPlanner( + 4, + 1, + min_avg_tokens_per_expert=0, + rebalance_gain_threshold=0.0, ) - selected, aligned_improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, - current, - raw_improving_candidate, - rebalance_gain_threshold=0.05, - expert_alignment=128, + current = _initial_extra_expert_placement(16, 4, 1).unsqueeze(0).tolist() + load = torch.randint( + 0, + 10000, + (1, 4, 16), + generator=torch.Generator().manual_seed(2), ) - assert raw_improved.item() - assert not aligned_improved.item() - assert torch.equal(selected, current) - + result = planner.plan(load.sum(dim=1).tolist(), current) -def test_estimate_rank_load_aligns_each_sample_before_accumulation(): - samples = torch.tensor([[[20, 0, 0, 0]], [[20, 0, 0, 0]]]).unsqueeze(2) - placement = torch.tensor([[[2], [0]]]) + assert len(result.placement) == len(current) + assert all(len(actual) == len(expected) for actual, expected in zip(result.placement[0], current[0])) + for rank, row in enumerate(result.placement[0]): + assert all(expert // 4 != rank for expert in row) - per_sample = _estimate_rank_load(samples, placement, expert_alignment=128) - accumulated = _estimate_rank_load(samples.sum(dim=0, keepdim=True), placement, expert_alignment=128) - assert torch.equal(per_sample[:, 0], torch.tensor([[128.0, 128.0], [128.0, 128.0]])) - assert torch.equal(per_sample.sum(dim=0)[0], torch.tensor([256.0, 256.0])) - assert torch.equal(accumulated[0, 0], torch.tensor([128.0, 128.0])) - - -def test_select_improving_placements_rejects_lower_ratio_when_critical_is_unchanged(): - samples = torch.tensor( - [ - [[255, 220, 226, 254]], - [[172, 278, 51, 238]], - [[249, 291, 284, 183]], - ] - ).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - mean_inflating_candidate = torch.tensor([[[2], [1]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - samples, - current, - mean_inflating_candidate, - rebalance_gain_threshold=0.05, - expert_alignment=128, - ) - current_load = _estimate_rank_load(samples, current, expert_alignment=128) - candidate_load = _estimate_rank_load(samples, mean_inflating_candidate, expert_alignment=128) - current_critical = current_load.max(dim=2).values.sum() - candidate_critical = candidate_load.max(dim=2).values.sum() - - assert candidate_load.mean(dim=2).sum() > current_load.mean(dim=2).sum() - assert current_critical == candidate_critical - assert not improved.item() - assert torch.equal(selected, current) - - -def test_select_improving_placements_accepts_five_percent_critical_reduction(): - samples = torch.tensor( - [ - [[13, 352, 348, 141]], - [[287, 175, 236, 179]], - [[316, 99, 266, 353]], - ] - ).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[2], [1]]]) +def test_eplb_planner_keeps_high_redundancy_search_state_isolated(): + planner = GreedyEPLBPlanner(4, 3, min_avg_tokens_per_expert=0) + current = _initial_extra_expert_placement(16, 4, 3).unsqueeze(0).tolist() + load = [ + [22613, 26852, 21852, 23480, 13270, 14695, 28735, 22303, 15324, 19604, 21492, 25458, 14120, 12130, 18620, 22888] + ] - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - samples, - current, - candidate, - rebalance_gain_threshold=0.05, - expert_alignment=128, - ) + result = planner.plan(load, current) - assert improved.item() - assert torch.equal(selected, candidate) + for rank, row in enumerate(result.placement[0]): + assert len(row) == len(set(row)) == 3 + assert all(expert // 4 != rank for expert in row) -def test_select_improving_placements_rejects_single_layer_gain_below_model_threshold(): - # Layer 0 becomes better, but layer 1 dominates model critical load. The - # aggregate estimated critical-load reduction gain is below 5%, so neither layer may be changed. - expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 10]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]], [[2], [0]]]) - candidate = torch.tensor([[[3], [0]], [[2], [0]]]) +def test_fused_moe_loads_default_replicas_into_their_physical_rows(): + weight = object.__new__(fused_weight_module.FusedMoeWeight) + weight.lock = threading.Lock() + loaded = [] - selected, improved, metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) + def load_weight(expert, local, _weights): + loaded.append(("weight", expert, local)) - assert not torch.any(improved) - assert torch.equal(selected, current) - assert metrics["candidate_rebalance_gain"] == pytest.approx(0.5 / 13) - assert metrics["candidate_changed_layer_count"] == 1 + def load_scale(expert, local, _weights): + loaded.append(("scale", expert, local)) + def load_zero_point(expert, local, _weights): + loaded.append(("zero", expert, local)) -def test_select_improving_placements_accepts_only_when_model_gain_reaches_threshold(): - expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 5]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]], [[2], [0]]]) - candidate = torch.tensor([[[3], [0]], [[2], [0]]]) + weight._load_expert = load_weight + weight._load_expert_scale = load_scale + weight._load_expert_zero_point = load_zero_point + local_logic_expert_ids_list = build_initial_local_expert_ids(8, 4, 2)[0] - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) + weight._load_weight(local_logic_expert_ids_list, {}) - assert torch.equal(improved, torch.tensor([True, False])) - assert torch.equal(selected, candidate) + assert local_logic_expert_ids_list == [0, 1, 2, 3] + assert [entry for entry in loaded if entry[0] == "weight"] == [ + ("weight", 0, 0), + ("weight", 1, 1), + ("weight", 2, 2), + ("weight", 3, 3), + ] def test_logical_to_physical_map_selects_one_physical_expert(): @@ -746,603 +540,17 @@ def test_logical_to_physical_maps_for_layers_match_single_layer_api(current_rank ) -def test_plan_redundant_experts_prefers_current_rank_load_relief(): - source_load = torch.zeros((1, 1, 4, 8), dtype=torch.int64) - source_load[0, 0, 1, 0] = 1024 - placement = plan_redundant_experts( - source_load, - num_ranks=4, - num_redundant_experts_per_rank=1, - expert_alignment=128, - ) - assert placement[0, 1, 0] == 0 - predicted = _estimate_rank_load(source_load, placement, expert_alignment=128) - assert torch.equal( - predicted, - _manual_runtime_rank_load(source_load, placement, alignment=128), - ) - - -def test_plan_redundant_experts_accepts_aggregated_load(): - load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]).unsqueeze(0).unsqueeze(2) - placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1) - assert placement.shape == (1, 4, 1) - - -def test_current_rank_estimate_matches_local_first_runtime_replica_sharing(): - placement = torch.tensor([[[4], [5], [0], [1]]], dtype=torch.int64) - source_load = torch.zeros((1, 1, 4, 8), dtype=torch.int64) - source_load[0, 0, 0, 0] = 256 - source_load[0, 0, 2, 0] = 128 - - predicted = _estimate_rank_load(source_load, placement, expert_alignment=128) - runtime = _manual_runtime_rank_load(source_load, placement, alignment=128) - - assert torch.equal(predicted, runtime) - assert torch.equal(predicted[0, 0], torch.tensor([256.0, 0.0, 128.0, 0.0])) - - -def test_current_rank_planner_constraints_and_real_critical_improvement(): - load_by_node = torch.tensor( - [ - [ - [ - [697, 451, 383, 536, 349, 404, 854, 425], - [861, 103, 166, 612, 444, 263, 910, 392], - ] - ], - [ - [ - [944, 457, 338, 108, 63, 525, 48, 216], - [439, 117, 837, 550, 833, 201, 729, 5], - ] - ], - [ - [ - [749, 159, 18, 723, 12, 700, 419, 51], - [112, 135, 8, 840, 40, 970, 90, 683], - ] - ], - ], - dtype=torch.int64, - ) - source_load = torch.zeros((3, 1, 4, 8), dtype=torch.int64) - source_load[:, :, 0] = load_by_node[:, :, 0] - source_load[:, :, 2] = load_by_node[:, :, 1] - initial = torch.tensor([[[2], [4], [6], [0]]], dtype=torch.int64) - planned = plan_redundant_experts(source_load, 4, 1, expert_alignment=128) - - for rank, experts in enumerate(planned[0].tolist()): - assert len(experts) == len(set(experts)) == 1 - assert experts[0] // 2 != rank - - before = _estimate_rank_load(source_load, initial, expert_alignment=128) - after = _estimate_rank_load(source_load, planned, expert_alignment=128) - manual_before = _manual_runtime_rank_load(source_load, initial, 128) - manual_after = _manual_runtime_rank_load(source_load, planned, 128) - assert torch.equal(before, manual_before) - assert torch.equal(after, manual_after) - assert after.max(dim=2).values.sum() < before.max(dim=2).values.sum() - - -def test_current_rank_select_uses_the_same_runtime_critical_prediction(): - source_load = torch.zeros((2, 1, 4, 8), dtype=torch.int64) - source_load[:, 0, 1, 0] = torch.tensor([1024, 768]) - source_load[:, 0, 0, 6] = torch.tensor([896, 1024]) - current = torch.tensor([[[2], [4], [6], [0]]], dtype=torch.int64) - candidate = plan_redundant_experts(source_load, 4, 1, expert_alignment=128) - selected, _improved, _metrics, _before_load, _after_load = select_improving_placements( - source_load, - current, - candidate, - rebalance_gain_threshold=0.05, - expert_alignment=128, - ) - assert torch.equal( - _estimate_rank_load(source_load, selected, 128), - _manual_runtime_rank_load(source_load, selected, 128), - ) - - -def _count_moved_slots(current: torch.Tensor, target: torch.Tensor) -> int: - """Count rank rows gaining an expert: one migrated row per new expert id.""" - moved = 0 - for layer in range(current.shape[0]): - for rank in range(current.shape[1]): - moved += len(set(target[layer, rank].tolist()) - set(current[layer, rank].tolist())) - return moved - - -def test_sticky_plan_reproduces_current_when_load_unchanged(): - generator = torch.Generator().manual_seed(7) - load = torch.randint(1, 1000, (3, 16, 32), generator=generator).unsqueeze(2) - placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - - replanned = plan_redundant_experts( - load, - num_ranks=4, - num_redundant_experts_per_rank=2, - current_placement=placement, - stickiness=0.1, - ) - - assert torch.equal(replanned, placement) - for layer in range(placement.shape[0]): - assert build_transfer_plan(placement[layer], replanned[layer], 32, 4, 4) == [] - - -def test_sticky_plan_bounded_moves_under_small_perturbation(): - generator = torch.Generator().manual_seed(11) - load = torch.randint(100, 1000, (4, 16, 32), generator=generator).unsqueeze(2) - placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - noise = torch.rand((4, 16, 32), generator=generator).unsqueeze(2) * 0.1 + 0.95 - perturbed = (load.double() * noise).round().to(torch.int64) - - sticky = plan_redundant_experts(perturbed, 4, 2, current_placement=placement, stickiness=0.1) - free = plan_redundant_experts(perturbed, 4, 2) - - sticky_moves = _count_moved_slots(placement, sticky) - free_moves = _count_moved_slots(placement, free) - assert sticky_moves <= placement.numel() // 4 - assert sticky_moves < free_moves - - def critical(candidate): - return _estimate_rank_load(perturbed, candidate).max(dim=2).values.sum() - - assert critical(sticky) <= critical(free) * 1.1 - - -def test_sticky_plan_still_churns_under_phase_shift(): - layers, experts = 8, 32 - before = torch.full((layers, experts), 10, dtype=torch.int64) - after = torch.full((layers, experts), 10, dtype=torch.int64) - offsets = torch.arange(4) - for layer in range(layers): - before[layer, (4 * layer + offsets) % experts] = 5000 - after[layer, (4 * layer + 16 + offsets) % experts] = 5000 - before = before.unsqueeze(0).unsqueeze(2) - after = after.unsqueeze(0).unsqueeze(2) - placement = plan_redundant_experts(before, num_ranks=4, num_redundant_experts_per_rank=2) - - replanned = plan_redundant_experts( - after, - num_ranks=4, - num_redundant_experts_per_rank=2, - current_placement=placement, - stickiness=0.1, - ) - - assert _count_moved_slots(placement, replanned) > placement.numel() // 2 - - -def test_transfer_plan_slot_permutation_is_free(): +def test_transfer_plan_respects_explicit_target_slots(): current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) target = torch.tensor([[5, 4], [7, 6], [1, 0], [3, 2]]) - assert torch.equal(align_target_placement(current, target), current) - assert build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) == [] - - -def test_align_target_placement_keeps_retained_experts_in_live_slots(): - current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) - target = torch.tensor([[5, 6], [7, 0], [1, 0], [3, 2]]) - - canonical = align_target_placement(current, target) - - assert torch.equal(canonical, torch.tensor([[6, 5], [0, 7], [0, 1], [2, 3]])) - - -def test_align_target_placement_replaces_duplicate_current_replicas(): - current = torch.tensor([[1, 1], [3, 3]]) - target = torch.tensor([[1, 2], [3, 0]]) - - canonical = align_target_placement(current, target) - - assert torch.equal(canonical, target) - assert build_transfer_plan(current, target, num_logical_experts=4, world_size=2, node_world_size=2) - - -def test_canonical_placement_keeps_transfer_rows_and_published_map_consistent(): - num_logical_experts = 8 - world_size = 4 - num_redundant_slots_per_rank = 2 - num_experts_per_rank = num_logical_experts // world_size - num_physical_experts_per_rank = num_experts_per_rank + num_redundant_slots_per_rank - current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) - target = torch.tensor([[5, 6], [7, 0], [1, 0], [3, 2]]) - canonical = align_target_placement(current, target) - plan = build_transfer_plan(current, canonical, num_logical_experts, world_size, node_world_size=2) - - # Label every current physical row by its resident logical expert, then - # apply the transfer plan from a frozen source snapshot just as staging - # copies do before the destination rows are published. - source_rows = [ - list(range(rank * num_experts_per_rank, (rank + 1) * num_experts_per_rank)) + current[rank].tolist() - for rank in range(world_size) - ] - live_rows = [row.copy() for row in source_rows] - for step in plan: - live_rows[step.dst_rank][num_experts_per_rank + step.dst_slot] = source_rows[step.src_rank][step.src_local_row] - - for current_rank in range(world_size): - logical_to_physical = build_logical_to_physical_map( - _rank_to_logic_expert_ids(canonical.tolist(), num_logical_experts), - num_logical_experts, - current_rank=current_rank, - ) - for logical_expert, packed_row in enumerate(logical_to_physical): - count = packed_row[0] - for physical_id in packed_row[2 : count + 2]: - rank, row = divmod(physical_id, num_physical_experts_per_rank) - assert live_rows[rank][row] == logical_expert - - -def test_plan_and_broadcast_publishes_canonical_placement(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.global_rank = 0 - manager.world_size = 4 - manager.node_world_size = 2 - manager.num_logical_experts = 8 - manager.num_redundant_experts_per_rank = 2 - manager.current_placement = torch.tensor([[[4, 5], [6, 7], [0, 1], [2, 3]]]) - manager.placement_stickiness = 0.1 - manager.rebalance_gain_threshold = 0.05 - manager.evaluation_group = object() - candidate = torch.tensor([[[5, 4], [7, 6], [1, 0], [3, 2]]]) - broadcasts = [] - - def fixed_selector(*_args, **_kwargs): - rank_load = torch.full((1, 1, 4), 100.0) - return candidate.clone(), torch.tensor([True]), {}, rank_load, rank_load - - def record_broadcast(result_list, **_kwargs): - broadcasts.append(result_list[0]) - - monkeypatch.setattr( - manager_module, - "plan_redundant_experts", - lambda *_args, **_kwargs: candidate.clone(), - ) - monkeypatch.setattr(manager_module, "select_improving_placements", fixed_selector) - monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) - - result = manager._plan_and_broadcast(torch.full((1, 1, 2, 8), 100, dtype=torch.int64)) - - assert torch.equal(result["placement"], manager.current_placement) - assert broadcasts and torch.equal(broadcasts[0]["placement"], manager.current_placement) - - -def test_stickiness_zero_matches_unbiased_plan(): - generator = torch.Generator().manual_seed(17) - load = torch.randint(1, 1000, (2, 8, 16), generator=generator).unsqueeze(2) - unbiased = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - unrelated = _initial_extra_expert_placement(16, 4, 2).unsqueeze(0).expand(8, -1, -1).clone() - - replanned = plan_redundant_experts( - load, - num_ranks=4, - num_redundant_experts_per_rank=2, - current_placement=unrelated, - stickiness=0.0, - ) - - assert torch.equal(replanned, unbiased) - - -def test_plan_and_broadcast_propagates_rank_zero_error_after_existing_broadcast( - monkeypatch, -): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.global_rank = 0 - manager.world_size = 2 - manager.node_world_size = 2 - manager.num_logical_experts = 8 - manager.num_redundant_experts_per_rank = 2 - manager.current_placement = torch.tensor([[[4, 5], [6, 7], [0, 1], [2, 3]]]) - manager.placement_stickiness = 0.1 - manager.rebalance_gain_threshold = 0.05 - manager.evaluation_group = object() - broadcasted = [] - - def broken_planner(*_args, **_kwargs): - raise RuntimeError("planner boom") - - def record_broadcast(result_list, **_kwargs): - broadcasted.append(result_list[0]) - - monkeypatch.setattr(manager_module, "plan_redundant_experts", broken_planner) - monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) - - with pytest.raises(RuntimeError, match="EPLB planner failed on rank zero") as exc_info: - manager._plan_and_broadcast(torch.full((1, 1, 1, 8), 100, dtype=torch.int64)) - - assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "planner boom" - assert broadcasted == [{"kind": "error", "message": "RuntimeError: planner boom"}] - - -def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_evaluates_step_twenty( - monkeypatch, -): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - counter = torch.tensor([10, 20, 30, 40], dtype=torch.int64) - manager.weights = [ - type( - "Weight", - (), - { - "fuse_moe_impl": _test_moe_impl( - eplb=True, - route_counter=counter, - recording=False, - num_logical_experts=4, - world_size=1, - ), - }, - )() - ] - manager.in_flight = False - manager.prefill_steps = 15 - manager.step_interval = 20 - manager.sampling_interval = manager.step_interval - manager.evaluation_group = object() - manager.num_logical_experts = 4 - manager.global_rank = 1 - manager.evaluation_in_flight = False - manager._steady_collection_end_step = None - manager._continuous_collection_start_step = None - manager._continuous_collection_end_step = None - - recordings, resets, started = [], [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_route_counters = lambda: resets.append(True) - monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) - monkeypatch.setattr( - manager_module.torch.cuda, - "synchronize", - lambda: pytest.fail("step must not synchronize CUDA"), - ) - - manager.step() - assert manager.prefill_steps == 16 - assert recordings == [True] - assert resets == [True] - assert manager._steady_collection_end_step == 20 - - for _ in range(3): - manager.step() - assert manager.prefill_steps == 19 - assert recordings == [True] - assert started == [] - - manager.step() - assert manager.prefill_steps == 20 - assert manager._steady_collection_end_step is None - assert started == [True] - - -def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary( - monkeypatch, -): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.evaluation_in_flight = False - manager.prefill_steps = 0 - manager.step_interval = 20 - manager.sampling_interval = 3 - manager._steady_collection_end_step = None - manager._continuous_collection_start_step = None - manager._continuous_collection_end_step = None - recordings, resets, started = [], [], [] - manager._set_recording = lambda enabled: recordings.append((manager.prefill_steps, enabled)) - manager._reset_route_counters = lambda: resets.append(manager.prefill_steps) - monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(manager.prefill_steps)) - - manager._prepare_next_sampling_window() - assert resets == [0] - assert recordings == [(0, True)] - assert manager._steady_collection_end_step == 3 - - manager.step() - manager.step() - assert started == [] - manager.step() - assert started == [3] - - -def test_eplb_step_does_not_start_a_second_evaluation_while_one_is_pending(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - counter = torch.tensor([10, 20, 30, 40], dtype=torch.int64) - manager.weights = [ - type( - "Weight", - (), - { - "fuse_moe_impl": _test_moe_impl( - eplb=True, - route_counter=counter, - recording=False, - num_logical_experts=4, - world_size=1, - ), - }, - )() - ] - manager.in_flight = False - manager.prefill_steps = 1 - manager.step_interval = 2 - manager.sampling_interval = manager.step_interval - manager.evaluation_group = object() - manager.num_logical_experts = 4 - manager.global_rank = 1 - manager.evaluation_in_flight = True - - started = [] - monkeypatch.setattr(manager, "_poll_evaluation", lambda: True) - monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) - - manager.step() - - assert manager.prefill_steps == 1 - assert started == [] - - -def test_evaluation_no_improvement_logs_model_fields_without_reopening_interval_window( - monkeypatch, -): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 2, - } - manager.global_rank = 0 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.weights = [] - manager._eplb_impls = [] - recordings, logs = [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) - - assert not manager._poll_evaluation() - assert recordings == [False] - assert "model_imbalance_ratio" in logs[0][0] - assert "candidate_rebalance_gain" in logs[0][0] - assert "candidate_changed_layer_count" in logs[0][0] - assert "actual_changed_layer_count" in logs[0][0] - assert "next_sampling_interval" in logs[0][0] - assert manager.sampling_interval == 80 - - -def test_interval_one_rearms_after_evaluation_but_never_evaluates_empty_counter( - monkeypatch, -): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.prefill_steps = 1 - manager.step_interval = 1 - manager.sampling_interval = 1 - manager._steady_collection_end_step = None - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 2, - } - manager.global_rank = 1 - manager.weights = [] - manager._eplb_impls = [] - recordings, starts = [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._start_evaluation = lambda: starts.append(True) - manager._evaluation_ready_on_all_ranks = lambda: True - - manager.poll() - - # The no-improvement backoff changes interval 1 to 4. The clamped - # steady window arms immediately but still waits for boundary step 5. - assert recordings == [True] - assert manager._steady_collection_end_step is not None - assert manager._steady_collection_end_step == 5 - assert starts == [] - assert manager.prefill_steps == 1 - - for _ in range(3): - manager.step() - assert starts == [] - manager.step() - assert starts == [True] - - -def test_evaluation_worker_error_is_raised_by_main_thread(): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = RuntimeError("planner failed") - manager._evaluation_result = None - manager._evaluation_thread = DoneThread() - - with pytest.raises(RuntimeError, match="planner failed"): - manager._poll_evaluation() - - -def test_evaluation_state_is_cleared_before_second_round(monkeypatch): - class DoneThread: - def join(self): - pass - - class PendingThread: - def __init__(self, **_kwargs): - self.started = False - - def start(self): - self.started = True - - def join(self): - pytest.fail("pending worker must not be joined") - - class Event: - def record(self, _stream): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 1, - } - manager.global_rank = 1 - manager.prefill_steps = 0 - manager.step_interval = 1 - manager.sampling_interval = 1 - manager.weights = [] - manager._eplb_impls = [] - manager._set_recording = lambda _enabled: None - monkeypatch.setattr(manager_module.threading, "Thread", PendingThread) - monkeypatch.setattr(manager_module.torch.cuda, "Event", Event) - monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: object()) - - assert not manager._poll_evaluation() - assert manager._evaluation_result is None - assert manager._evaluation_error is None + plan = build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) - manager._start_evaluation() - assert manager.evaluation_in_flight - assert manager._evaluation_result is None - assert manager._poll_evaluation() # New worker has not produced a result. + assert len(plan) == target.numel() + assert {(step.dst_rank, step.dst_slot) for step in plan} == {(rank, slot) for rank in range(4) for slot in range(2)} -def test_manager_collects_aggregated_route_counters(): +def test_manager_collects_logical_route_counters(): counters = [ torch.tensor([10, 11], dtype=torch.int64), torch.tensor([40, 41], dtype=torch.int64), @@ -1367,21 +575,42 @@ def test_manager_collects_aggregated_route_counters(): manager._eplb_impls = [weight.fuse_moe_impl for weight in manager.weights] manager.num_logical_experts = 2 - samples = manager._collect_local_samples() + samples = manager._collect_local_load() assert torch.equal( samples, - torch.tensor([[[10, 11], [40, 41]]], dtype=torch.int64), + torch.tensor([[10, 11], [40, 41]], dtype=torch.int64), ) +def test_manager_delegates_distribution_planning_to_planner_class(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.global_rank = 0 + manager.world_size = 1 + manager.evaluation_group = object() + manager.current_placement = torch.tensor([[[1]]]) + logical_load = torch.tensor([[10, 20]]) + calls = [] + + class Result: + def as_dict(self): + return {"kind": "no_improvement"} + + manager.planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or Result())) + + result = manager._plan_and_broadcast(logical_load) + + assert result == {"kind": "no_improvement"} + assert len(calls) == 1 + assert calls[0][0] == logical_load.tolist() + assert calls[0][1] == manager.current_placement.tolist() + + def test_expert_load_imbalance_ratio_averages_layer_ratios(): global_load = torch.tensor( [ - [ - [[1, 2, 3], [1, 2, 3]], - [[4, 5, 6], [6, 5, 4]], - ] + [2, 4, 6], + [10, 10, 10], ], dtype=torch.int64, ) @@ -1397,7 +626,7 @@ def test_manager_publishes_expert_load_metrics_from_rank_zero(): manager.global_rank = 0 manager.metric_client = SimpleNamespace(gauge_set=lambda name, value: calls.append((name, value))) - manager._publish_expert_load_metrics( + manager._publish_expert_load_metric( { "expert_imbalance_ratio": 1.25, } @@ -1472,7 +701,7 @@ def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): assert not hasattr(impl, "recording") -def test_manager_evaluation_collective_preserves_current_rank_axis(monkeypatch): +def test_manager_evaluation_collective_sums_rank_loads(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.weights = [ type( @@ -1493,7 +722,6 @@ def test_manager_evaluation_collective_preserves_current_rank_axis(monkeypatch): manager.world_size = 4 manager.node_world_size = 2 manager.step_interval = 20 - manager.sampling_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0) @@ -1501,16 +729,15 @@ def test_manager_evaluation_collective_preserves_current_rank_axis(monkeypatch): manager._evaluation_lock = threading.Lock() manager._evaluation_result = None manager._evaluation_error = None - manager._continuous_collection_end_step = None - local = torch.full((1, 1, 4), 100, dtype=torch.int64) - manager._collect_local_samples = lambda: local + local = torch.full((1, 4), 100, dtype=torch.int64) + manager._collect_local_load = lambda: local.clone() seen = {} def all_reduce(tensor, **kwargs): seen["before"] = tensor.clone() seen["group"] = kwargs["group"] - # Simulate current rank 0's contribution from another process. - tensor[:, :, 0].fill_(100) + # Simulate one other rank contributing the same logical-expert load. + tensor.add_(100) monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) @@ -1524,19 +751,11 @@ def plan_and_broadcast(global_load): manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) assert seen["group"] is manager.evaluation_group - expected_local = local - assert seen["before"].shape == (1, 1, 4, 4) - assert torch.equal(seen["before"][:, :, 0], torch.zeros_like(expected_local)) - assert torch.equal(seen["before"][:, :, 1], torch.zeros_like(expected_local)) - assert torch.equal(seen["before"][:, :, 2], expected_local) - assert torch.equal(seen["before"][:, :, 3], torch.zeros_like(expected_local)) - assert torch.equal(seen["global_load"][:, :, 0], torch.full_like(expected_local, 100)) - assert torch.equal(seen["global_load"][:, :, 1], torch.zeros_like(expected_local)) - assert torch.equal(seen["global_load"][:, :, 2], expected_local) - assert torch.equal(seen["global_load"][:, :, 3], torch.zeros_like(expected_local)) + assert seen["before"].shape == (1, 4) + assert torch.equal(seen["before"], local) + assert torch.equal(seen["global_load"], torch.full_like(local, 200)) assert manager._evaluation_error is None assert manager._evaluation_result["expert_imbalance_ratio"] == 1.0 - assert manager._evaluation_result["sample_window_steps"] == 4 def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_call( @@ -1563,7 +782,6 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c manager.world_size = 4 manager.node_world_size = 2 manager.step_interval = 20 - manager.sampling_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() @@ -1571,8 +789,7 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c manager._evaluation_lock = threading.Lock() manager._evaluation_result = None manager._evaluation_error = None - manager._continuous_collection_end_step = None - manager._collect_local_samples = lambda: torch.full((1, 3, 4), 100, dtype=torch.int64) + manager._collect_local_load = lambda: torch.full((3, 4), 100, dtype=torch.int64) planned_placement = torch.tensor( [ [[3], [0], [1], [2]], @@ -1583,8 +800,8 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c ) manager._plan_and_broadcast = lambda _global_load: { "kind": "planned", - "placement": planned_placement, - "improved": torch.tensor([True, False, True]), + "placement": planned_placement.tolist(), + "changed_layers": [True, False, True], } monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) @@ -2077,7 +1294,9 @@ def narrow(self, _dim, start, length): assert copies == [("live", 11, 3, "staging", 1, 3)] -def test_manager_inflight_ready_gate_commits_one_layer_and_propagates_worker_error(monkeypatch): +def test_manager_inflight_ready_gate_commits_one_layer_and_propagates_worker_error( + monkeypatch, +): class Transfer: def __init__(self): self.ready = 0 @@ -2284,365 +1503,52 @@ def finish(self): g_infer_context.overlap_stream = original_overlap_stream -def test_manager_rearms_after_rebalance_for_interval_one(): - recording_calls = [] - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.step_interval = 1 - manager.sampling_interval = 1 - manager._steady_collection_end_step = None - manager._continuous_collection_start_step = 0 - manager.weights = [] - manager._eplb_impls = [] - manager.target_placement = torch.zeros(1) - manager.in_flight_started_at = 0 - manager.global_rank = 1 - manager._set_recording = lambda enabled: recording_calls.append(enabled) - manager._finish_rebalance() - assert manager.in_flight is False - assert recording_calls == [True] - assert manager._steady_collection_end_step is None - assert manager._continuous_collection_start_step is None - - -def test_manager_sparse_insufficient_schedules_bounded_fresh_window(monkeypatch): - class DoneThread: - def join(self): - pass - +def test_manager_inflight_step_does_not_poll(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "insufficient", - "minimum_layer_samples": 1, - "minimum": 2, - } - manager.global_rank = 0 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.prefill_steps = 37 - manager.weights = [] - manager._eplb_impls = [] - manager._continuous_collection_start_step = None - manager._continuous_collection_end_step = None - recordings, resets, logs = [], [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_route_counters = lambda: resets.append(True) - monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) - - assert not manager._poll_evaluation() - assert not hasattr(manager, "_retained_local_samples") - assert manager._continuous_collection_start_step == 40 - assert manager._continuous_collection_end_step == 60 - assert recordings == [False] - assert resets == [True] - assert manager.sampling_interval == 20 - assert "insufficient samples" in logs[0][0] - assert "scheduled_fresh_window" in logs[0][0] - - manager.in_flight = False - manager.evaluation_in_flight = False - starts = [] - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(manager.prefill_steps)) - manager.step() - manager.step() - assert manager.prefill_steps == 39 - assert starts == [] - manager.step() - assert manager.prefill_steps == 40 - assert starts == [] - for _ in range(19): - manager.step() - assert manager.prefill_steps == 59 - assert starts == [] + manager.in_flight = True + calls = [] + manager._poll_in_flight = lambda: calls.append("poll") manager.step() - assert starts == [60] - + assert calls == [] -def test_manager_full_window_insufficient_clears_and_backs_off(): - class DoneThread: - def join(self): - pass - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "insufficient", - "minimum_layer_samples": 1, - "minimum": 2, - } - manager.global_rank = 1 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.prefill_steps = 60 - manager.weights = [] - manager._eplb_impls = [] - manager._continuous_collection_start_step = 40 - manager._continuous_collection_end_step = 60 - recordings, resets = [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_route_counters = lambda: resets.append(True) - - assert not manager._poll_evaluation() - assert not hasattr(manager, "_retained_local_samples") - assert manager._continuous_collection_start_step is None - assert manager._continuous_collection_end_step is None - assert manager.sampling_interval == 80 - assert recordings == [False] - assert resets == [True] - - -def test_begin_continuous_collection_uses_full_window_at_fixed_boundary(monkeypatch): +def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.in_flight = False manager.evaluation_in_flight = False - manager.prefill_steps = 36 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager._steady_collection_end_step = manager.prefill_steps + 1 - recordings, resets, starts = [], [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_route_counters = lambda: resets.append(True) - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) - - # The window is never truncated to the next boundary: it waits until 40, - # records a full 20 fresh steps, then evaluates at the boundary at 60. - manager._begin_continuous_collection() - assert manager._continuous_collection_start_step == 40 - assert manager._continuous_collection_end_step == 60 - assert recordings == [False] - assert resets == [True] - assert manager._steady_collection_end_step is None - - for _ in range(4): - manager.step() - assert manager.prefill_steps == 40 - assert recordings == [False, True] - assert starts == [] - for _ in range(19): - manager.step() - assert starts == [] - manager.step() # 60: the full window ends and triggers the evaluation. - assert starts == [True] - - -def test_begin_continuous_collection_preserves_full_window_at_sparse_boundary(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.prefill_steps = 80 - manager.step_interval = 20 - manager.sampling_interval = 80 - manager._steady_collection_end_step = None - recordings = [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_route_counters = lambda: None - - manager._begin_continuous_collection() - assert manager._continuous_collection_start_step == 140 - assert manager._continuous_collection_end_step == 160 - assert recordings == [False] - - -def test_first_no_improvement_switches_to_sparse_sampling_window(monkeypatch): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 2, - } - manager.global_rank = 1 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager._continuous_collection_start_step = 0 - manager.weights = [] - manager._eplb_impls = [] - recordings = [] - manager._set_recording = lambda enabled: recordings.append(enabled) - - assert not manager._poll_evaluation() - assert manager._continuous_collection_start_step is None - assert recordings == [False] - assert manager.sampling_interval == 80 - - -def test_continuous_collection_evaluates_only_after_one_full_base_window(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager._continuous_collection_start_step = 0 - manager._continuous_collection_end_step = 20 manager.prefill_steps = 0 - manager.step_interval = 20 - manager.sampling_interval = 320 - manager._steady_collection_end_step = None - manager.evaluation_in_flight = False - started = [] - monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) - - for _ in range(19): - manager.step() - assert manager.prefill_steps == 19 - assert started == [] - - manager.step() - assert manager.prefill_steps == 20 - assert started == [True] - - -def test_no_improvement_exponentially_backs_off_sampling_interval_at_cap(): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager._evaluation_lock = threading.Lock() - manager._set_recording = lambda _enabled: None - manager.global_rank = 1 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.weights = [] - manager._eplb_impls = [] - - for expected_interval in (80, 320, 320): - manager.evaluation_in_flight = True - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 2, - } - assert not manager._poll_evaluation() - assert manager.sampling_interval == expected_interval - - -def test_sparse_backoff_arms_and_evaluates_only_at_new_interval_boundary(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.prefill_steps = 18 - manager.step_interval = 20 - manager.sampling_interval = 80 - manager._steady_collection_end_step = None - manager._continuous_collection_start_step = None - manager._continuous_collection_end_step = None - manager.evaluation_in_flight = False - recordings, resets, starts = [], [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_route_counters = lambda: resets.append(True) - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) + manager.step_interval = 3 + manager.next_evaluation_step = 3 + starts = [] + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(manager.prefill_steps)) manager.step() manager.step() - assert manager.prefill_steps == 20 - assert recordings == [] - assert starts == [] - - for _ in range(55): - manager.step() - assert manager.prefill_steps == 75 - assert recordings == [] assert starts == [] - manager.step() - assert manager.prefill_steps == 76 - assert recordings == [True] - assert resets == [True] - assert manager._steady_collection_end_step is not None - for _ in range(4): - manager.step() - assert manager.prefill_steps == 80 - assert starts == [True] - assert manager._steady_collection_end_step is None + assert starts == [3] -def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): +def test_manager_restart_collection_resets_counters_and_arms_next_window(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.current_placement = torch.zeros((1, 1, 1), dtype=torch.int64) - manager.num_logical_experts = 1 - manager.world_size = 1 - manager.node_world_size = 1 + manager.prefill_steps = 11 manager.step_interval = 20 - manager.sampling_interval = 320 - manager._continuous_collection_start_step = 0 - manager.global_rank = 1 - manager.transfer = type( - "Transfer", - (), - {"start": lambda self, plans: setattr(self, "started", plans)}, - )() - manager._reset_route_counters = lambda: None - - manager._start_rebalance( - { - "placement": torch.zeros((1, 1, 1), dtype=torch.int64), - "improved": torch.tensor([True]), - "metadata": [None], - "layer_plans": [(0, object())], - "before": {"max": 1.0, "p95": 1.0}, - "after": {"max": 1.0, "p95": 1.0}, - "model_imbalance_ratio": 1.0, - "candidate_model_imbalance_ratio": 1.0, - "candidate_rebalance_gain": 0.1, - "candidate_changed_layer_count": 1, - } + counter = torch.ones(4, dtype=torch.int64) + impl = _test_moe_impl( + eplb=True, + route_counter=counter, + recording=False, + num_logical_experts=4, + world_size=1, ) + manager._eplb_impls = [impl] - assert manager.sampling_interval == 20 - assert manager.in_flight - assert manager._continuous_collection_start_step is None - assert len(manager.transfer.started) == 1 - - -def test_first_rebalance_completion_switches_to_four_step_sparse_window(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.step_interval = 20 - manager.sampling_interval = manager.step_interval - manager._steady_collection_end_step = None - manager.weights = [] - manager._eplb_impls = [] - manager.target_placement = torch.zeros(1) - manager.in_flight_started_at = 0 - manager.global_rank = 1 - recordings, starts = [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._finish_rebalance() - assert recordings == [False] - - manager.in_flight = False - manager.prefill_steps = 38 - manager.evaluation_in_flight = False - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) - manager.prefill_steps = 35 - manager.step() - assert recordings == [False, True] - for _ in range(4): - manager.step() - assert starts == [True] - + manager._restart_collection() -def test_manager_inflight_step_does_not_poll(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = True - calls = [] - manager._poll_in_flight = lambda: calls.append("poll") - manager.step() - assert calls == [] + assert torch.count_nonzero(counter) == 0 + assert impl.recording + assert manager.next_evaluation_step == 31 def test_manager_poll_advances_inflight_before_evaluation(): @@ -2651,7 +1557,7 @@ def test_manager_poll_advances_inflight_before_evaluation(): manager.evaluation_in_flight = True calls = [] manager._poll_in_flight = lambda: calls.append("inflight") - manager._poll_evaluation = lambda: calls.append("evaluation") + manager._finish_evaluation = lambda: calls.append("evaluation") manager.poll() @@ -2668,7 +1574,7 @@ def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): manager.control_group = object() manager._control_ready_count = torch.empty(1, dtype=torch.int32) calls = [] - manager._poll_evaluation = lambda: calls.append("evaluation") + manager._finish_evaluation = lambda: calls.append("evaluation") monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(0)) manager.poll() @@ -2759,7 +1665,9 @@ def synchronize(self): assert transfer._copy_stream.synchronize_count == 2 -def test_pinned_transfer_waits_for_each_layer_commit_before_reusing_staging(monkeypatch): +def test_pinned_transfer_waits_for_each_layer_commit_before_reusing_staging( + monkeypatch, +): class Event: def __init__(self): self.synchronize_count = 0 @@ -2854,10 +1762,9 @@ def new_group(*args, **kwargs): ) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 assert transfer_calls == [([weight], groups[2], 0)] - assert manager.rebalance_gain_threshold == 0.07 - assert "rebalance_gain_threshold=0.0700" in logs[0] - assert manager._continuous_collection_start_step is None - assert manager._continuous_collection_end_step == manager.step_interval + assert manager.planner.rebalance_gain_threshold == 0.07 + assert manager.next_evaluation_step == manager.step_interval + assert "planner=GreedyEPLBPlanner" in logs[0] assert weight.fuse_moe_impl.recording assert manager._eplb_impls[0] is weight.fuse_moe_impl assert not hasattr(weight.fuse_moe_impl, "update_logical_expert_counter") From 5c86700496c6e12766ea05a09de9e7ba4749c370 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 03:11:58 +0000 Subject: [PATCH 23/72] fix --- .../server/router/model_infer/mode_backend/base_backend.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 23452426dc..6294beb155 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -356,8 +356,7 @@ def init_mtp_draft_model(self, main_kvargs: dict): model_cfg=draft_model_cfg, spec_mode=spec_mode, ) - draft_model = draft_model_class(draft_model_kvargs) - self.draft_models.append(draft_model) + self.draft_models.append(draft_model_class(draft_model_kvargs)) self.logger.info(f"loaded speculative draft model class {self.draft_models[i].__class__}") return From dc48e828e54ebcb37b521e2b4bc22208f0bdefeb Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 04:56:20 +0000 Subject: [PATCH 24/72] refactor(eplb): remove planner sample threshold --- .../meta_weights/fused_moe/eplb_planner.py | 28 ++----------------- .../model_infer/mode_backend/eplb_manager.py | 9 +----- unit_tests/common/fused_moe/test_eplb.py | 19 +++++-------- 3 files changed, 10 insertions(+), 46 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py index 84f5ac9625..bcf43a6b14 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py @@ -28,8 +28,6 @@ class EPLBPlan: before_rank_load: RankLoad after_rank_load: RankLoad reason: str - minimum_layer_samples: float - required_layer_samples: float @property def changed(self) -> bool: @@ -45,8 +43,6 @@ def as_dict(self) -> Dict[str, Any]: "kind": self.reason, "placement": self.placement, "changed_layers": self.changed_layers, - "minimum_layer_samples": self.minimum_layer_samples, - "required_layer_samples": self.required_layer_samples, "before": before, "after": after, "rebalance_gain": gain, @@ -75,7 +71,6 @@ def __init__( num_redundant_experts_per_rank: int, *, expert_alignment: int = 1, - min_avg_tokens_per_expert: float = 0, rebalance_gain_threshold: float = 0.0, ): if world_size <= 1: @@ -84,14 +79,11 @@ def __init__( raise ValueError("num_redundant_experts_per_rank must be positive") if expert_alignment <= 0: raise ValueError("expert_alignment must be positive") - if min_avg_tokens_per_expert < 0: - raise ValueError("min_avg_tokens_per_expert must be non-negative") if not 0.0 <= rebalance_gain_threshold <= 1.0: raise ValueError("rebalance_gain_threshold must be between 0.0 and 1.0") self.world_size = world_size self.num_redundant_experts_per_rank = num_redundant_experts_per_rank self.expert_alignment = expert_alignment - self.min_avg_tokens_per_expert = min_avg_tokens_per_expert self.rebalance_gain_threshold = rebalance_gain_threshold def plan( @@ -109,23 +101,9 @@ def plan( """ load = [[float(value) for value in layer] for layer in logical_expert_load] current = [[[int(expert) for expert in rank] for rank in layer] for layer in current_placement] - num_logical_experts = self._validate_inputs(load, current) - layer_samples = [sum(layer) for layer in load] - required_samples = self.min_avg_tokens_per_expert * num_logical_experts - minimum_samples = min(layer_samples) + self._validate_inputs(load, current) before_rank_load = self.estimate_rank_load(load, current) - if minimum_samples < required_samples: - return EPLBPlan( - placement=current, - changed_layers=[False] * len(load), - before_rank_load=before_rank_load, - after_rank_load=[layer[:] for layer in before_rank_load], - reason="insufficient", - minimum_layer_samples=minimum_samples, - required_layer_samples=required_samples, - ) - candidates = [self._plan_layer(layer_load, current_layer) for layer_load, current_layer in zip(load, current)] candidate_rank_load = self.estimate_rank_load(load, candidates) changed_layers = [] @@ -135,7 +113,7 @@ def plan( before = max(before_rank_load[layer]) after = max(candidate_rank_load[layer]) gain = (before - after) / max(before, 1.0) - changed = candidate != current[layer] and gain >= self.rebalance_gain_threshold + changed = candidate != current[layer] and gain > self.rebalance_gain_threshold changed_layers.append(changed) placement.append(candidate if changed else current[layer]) after_rank_load.append(candidate_rank_load[layer][:] if changed else before_rank_load[layer][:]) @@ -146,8 +124,6 @@ def plan( before_rank_load=before_rank_load, after_rank_load=after_rank_load, reason="planned" if any(changed_layers) else "no_improvement", - minimum_layer_samples=minimum_samples, - required_layer_samples=required_samples, ) def estimate_rank_load( diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 1c5ba9deda..3601f8b70c 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -34,7 +34,6 @@ from lightllm.utils.shm_port_args import get_shm_port_args logger = init_logger(__name__) -EPLB_MIN_AVG_TOKENS_PER_EXPERT = 100 EPLB_EXPERT_ALIGNMENT = 128 EPLB_CONTROL_ERROR = -1 EPLB_EXPERT_IMBALANCE_RATIO_METRIC = "lightllm_eplb_topk_expert_imbalance_ratio" @@ -77,7 +76,6 @@ def __init__(self, model: TpPartBaseModel): self.world_size, self.num_redundant_experts_per_rank, expert_alignment=EPLB_EXPERT_ALIGNMENT, - min_avg_tokens_per_expert=EPLB_MIN_AVG_TOKENS_PER_EXPERT, rebalance_gain_threshold=get_eplb_rebalance_gain_threshold(), ) @@ -268,12 +266,7 @@ def _finish_evaluation(self): self._publish_expert_load_metric(result) if result["kind"] != "planned": if self.global_rank == 0: - logger.info( - "eplb skip rearrangement kind=%s minimum_layer_samples=%s required_layer_samples=%s", - result["kind"], - result["minimum_layer_samples"], - result["required_layer_samples"], - ) + logger.info("eplb skip rearrangement kind=%s", result["kind"]) self._restart_collection() return self._start_rebalance(result) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 853d46f466..faa6ca46b4 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -294,7 +294,6 @@ def test_eplb_planner_builds_legal_concrete_slot_layout(): 4, 1, expert_alignment=1, - min_avg_tokens_per_expert=0, rebalance_gain_threshold=0.0, ) current = _initial_extra_expert_placement(8, 4, 1).unsqueeze(0).tolist() @@ -318,7 +317,6 @@ def test_eplb_planner_estimator_distributes_global_load_across_copies(): 4, 1, expert_alignment=128, - min_avg_tokens_per_expert=0, ) placement = [[[2], [4], [6], [0]]] load = [[100, 200, 300, 400, 500, 600, 700, 800]] @@ -328,15 +326,13 @@ def test_eplb_planner_estimator_distributes_global_load_across_copies(): assert predicted == [[640, 1024, 1280, 1408]] -def test_eplb_planner_rejects_an_under_sampled_window(): - planner = GreedyEPLBPlanner(2, 1, min_avg_tokens_per_expert=100) - current = [[[2], [0]]] - load = [[1, 1, 1, 1]] +def test_eplb_planner_does_not_move_zero_load_experts(): + planner = GreedyEPLBPlanner(2, 1) + current = [[[3], [1]]] - result = planner.plan(load, current) + result = planner.plan([[0, 0, 0, 0]], current) - assert result.reason == "insufficient" - assert not result.changed + assert result.reason == "no_improvement" assert result.placement == current @@ -344,7 +340,6 @@ def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): planner = GreedyEPLBPlanner( 4, 1, - min_avg_tokens_per_expert=0, rebalance_gain_threshold=0.0, ) current = _initial_extra_expert_placement(16, 4, 1).unsqueeze(0).tolist() @@ -364,7 +359,7 @@ def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): def test_eplb_planner_keeps_high_redundancy_search_state_isolated(): - planner = GreedyEPLBPlanner(4, 3, min_avg_tokens_per_expert=0) + planner = GreedyEPLBPlanner(4, 3) current = _initial_extra_expert_placement(16, 4, 3).unsqueeze(0).tolist() load = [ [22613, 26852, 21852, 23480, 13270, 14695, 28735, 22303, 15324, 19604, 21492, 25458, 14120, 12130, 18620, 22888] @@ -744,7 +739,7 @@ def all_reduce(tensor, **kwargs): def plan_and_broadcast(global_load): seen["global_load"] = global_load.clone() - return {"kind": "insufficient"} + return {"kind": "no_improvement"} manager._plan_and_broadcast = plan_and_broadcast From 9235c1816ac981ff440f2347b06caa09f1bf5938 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 05:01:09 +0000 Subject: [PATCH 25/72] fix(moe): restore random routing during autotune warmup --- .../basemodel/triton_kernel/fused_moe/topk_select.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py index d2f59de480..87cda6ce16 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py @@ -21,6 +21,7 @@ from lightllm.utils.sgl_utils import sgl_ops from typing import Callable, List, Optional, Tuple from lightllm.common.basemodel.triton_kernel.fused_moe.softmax_topk import softmax_topk +from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType def fused_topk( @@ -167,4 +168,12 @@ def select_experts( hidden_states=hidden_states, gating_output=router_logits, topk=top_k, renormalize=renormalize ) + ######################################## warning ################################################## + # here is used to match autotune feature, make topk_ids more random + if Autotuner.is_kernel_autotune_warmup(AutotuneKernelType.GENERAL): + rand_gen = torch.Generator(device="cuda") + rand_gen.manual_seed(router_logits.shape[0]) + router_logits = torch.randn(size=router_logits.shape, generator=rand_gen, dtype=torch.float32, device="cuda") + _, topk_ids = torch.topk(router_logits, k=top_k, dim=1) + return topk_weights, topk_ids From 99f055062fd7d2ad75e639dd69b49f4f2680c127 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 05:14:49 +0000 Subject: [PATCH 26/72] refactor(eplb): simplify manager state --- .../model_infer/mode_backend/eplb_manager.py | 97 +++++++------------ unit_tests/common/fused_moe/test_eplb.py | 60 ++++++------ 2 files changed, 62 insertions(+), 95 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 3601f8b70c..6601fe7e84 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -1,3 +1,4 @@ +from concurrent.futures import Future import threading import time @@ -43,12 +44,12 @@ class EPLBManager: """Collect logical load, invoke a planner, and publish migrated rows.""" def __init__(self, model: TpPartBaseModel): - self.weights = _find_fused_moe_weights(model) - assert self.weights, "EPLB requires at least one EP MoE layer" + weights = _find_fused_moe_weights(model) + assert weights, "EPLB requires at least one EP MoE layer" self.global_rank = get_global_rank() self.world_size = get_global_world_size() self.node_world_size = get_node_world_size() - self._eplb_impls = [weight.fuse_moe_impl for weight in self.weights] + self._eplb_impls = [weight.fuse_moe_impl for weight in weights] routed = {impl.n_routed_experts for impl in self._eplb_impls} redundant = {impl.num_redundant_experts_per_rank for impl in self._eplb_impls} @@ -69,7 +70,7 @@ def __init__(self, model: TpPartBaseModel): expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids ] self.current_placement = torch.tensor( - [initial_redundant_expert_ids for _ in self.weights], + [initial_redundant_expert_ids for _ in weights], dtype=torch.int64, ) self.planner: EPLBPlanner = GreedyEPLBPlanner( @@ -79,58 +80,45 @@ def __init__(self, model: TpPartBaseModel): rebalance_gain_threshold=get_eplb_rebalance_gain_threshold(), ) - self.in_flight = False self.in_flight_layers = [] - self.in_flight_started_at = None self.target_placement = None self.target_metadata = None - self.evaluation_in_flight = False - self._evaluation_lock = threading.Lock() - self._evaluation_result = None - self._evaluation_error = None - self._evaluation_thread = None + self._evaluation = None self.metric_client = None self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") self._control_ready_count = torch.empty(1, dtype=torch.int32) - self.transfer = PinnedMemoryEPLBTransfer(self.weights, self.transfer_group, self.global_rank) + self.transfer = PinnedMemoryEPLBTransfer(weights, self.transfer_group, self.global_rank) self._restart_collection() if self.global_rank == 0: logger.info( - f"eplb enabled layers={len(self.weights)} num_logical_experts={self.num_logical_experts} " + f"eplb enabled layers={len(weights)} num_logical_experts={self.num_logical_experts} " f"num_redundant_experts_per_rank={self.num_redundant_experts_per_rank} " f"step_interval={self.step_interval} planner={type(self.planner).__name__}" ) def poll(self): """Advance EPLB only from a globally ordered pre-forward boundary.""" - if self.in_flight: + if self.in_flight_layers: self._poll_in_flight() - elif self.evaluation_in_flight and self._evaluation_ready_on_all_ranks(): + elif self._evaluation is not None and self._evaluation_ready_on_all_ranks(): self._finish_evaluation() def step(self): - if self.in_flight or self.evaluation_in_flight: + if self.in_flight_layers or self._evaluation is not None: return self.prefill_steps += 1 if self.prefill_steps >= self.next_evaluation_step: self._start_evaluation() - def _set_recording(self, enabled: bool): - for impl in self._eplb_impls: - impl.recording = enabled - - def _reset_route_counters(self): - counters = [impl.route_counter for impl in self._eplb_impls] - if counters: - torch._foreach_zero_(counters) - def _restart_collection(self): - self._reset_route_counters() - self._set_recording(True) + counters = [impl.route_counter for impl in self._eplb_impls] + torch._foreach_zero_(counters) + for impl in self._eplb_impls: + impl.recording = True self.next_evaluation_step = self.prefill_steps + self.step_interval def _collect_local_load(self) -> torch.Tensor: @@ -139,9 +127,6 @@ def _collect_local_load(self) -> torch.Tensor: raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") return torch.stack(counters).cpu() - def _control_count(self, value: int) -> torch.Tensor: - return self._control_ready_count.fill_(value) - def _plan_and_broadcast(self, global_load: torch.Tensor): result = None local_error = None @@ -165,11 +150,9 @@ def _plan_and_broadcast(self, global_load: torch.Tensor): return result def _build_rebalance_data(self, result): - metadata = [None] * len(self.weights) + metadata = {} layer_plans = [] changed_layer_indices = [layer_index for layer_index, changed in enumerate(result["changed_layers"]) if changed] - if not changed_layer_indices: - return metadata, layer_plans num_primary_experts_per_rank = self.num_logical_experts // self.world_size rank_to_logic_expert_ids_by_layer = [] @@ -211,7 +194,7 @@ def _build_rebalance_data(self, result): ) return metadata, layer_plans - def _evaluate_after_event(self, event: torch.cuda.Event): + def _evaluate_after_event(self, event: torch.cuda.Event, evaluation: Future): try: torch.cuda.set_device(self._eplb_impls[0].route_counter.device) event.synchronize() @@ -221,29 +204,27 @@ def _evaluate_after_event(self, event: torch.cuda.Event): result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) if result["kind"] == "planned": result["metadata"], result["layer_plans"] = self._build_rebalance_data(result) - with self._evaluation_lock: - self._evaluation_result = result + evaluation.set_result(result) except BaseException as exc: - with self._evaluation_lock: - self._evaluation_error = exc + evaluation.set_exception(exc) def _start_evaluation(self): - with self._evaluation_lock: - self._evaluation_result = None - self._evaluation_error = None - self._set_recording(False) + for impl in self._eplb_impls: + impl.recording = False event = torch.cuda.Event() event.record(torch.cuda.current_stream()) - self.evaluation_in_flight = True - self._evaluation_thread = threading.Thread(target=self._evaluate_after_event, args=(event,), daemon=True) - self._evaluation_thread.start() + self._evaluation = Future() + threading.Thread( + target=self._evaluate_after_event, + args=(event, self._evaluation), + daemon=True, + ).start() def _evaluation_ready_on_all_ranks(self) -> bool: - with self._evaluation_lock: - error = self._evaluation_error - result = self._evaluation_result - status = EPLB_CONTROL_ERROR if error is not None else int(result is not None) - ready_count = self._control_count(status) + ready = self._evaluation.done() + error = self._evaluation.exception() if ready else None + status = EPLB_CONTROL_ERROR if error is not None else int(ready) + ready_count = self._control_ready_count.fill_(status) dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) ready_count = int(ready_count.item()) if ready_count < 0: @@ -253,16 +234,8 @@ def _evaluation_ready_on_all_ranks(self) -> bool: return bool(ready_count) def _finish_evaluation(self): - with self._evaluation_lock: - error = self._evaluation_error - result = self._evaluation_result - self._evaluation_error = None - self._evaluation_result = None - self._evaluation_thread.join() - self._evaluation_thread = None - self.evaluation_in_flight = False - if error is not None: - raise error + result = self._evaluation.result() + self._evaluation = None self._publish_expert_load_metric(result) if result["kind"] != "planned": if self.global_rank == 0: @@ -282,7 +255,6 @@ def _start_rebalance(self, result): self.target_placement = torch.tensor(result["placement"], dtype=torch.int64) self.target_metadata = result["metadata"] self.in_flight_layers = [layer_index for layer_index, _ in result["layer_plans"]] - self.in_flight = True self.in_flight_started_at = time.time() self.transfer.start(result["layer_plans"]) if self.global_rank == 0: @@ -308,7 +280,7 @@ def _poll_in_flight(self): except BaseException as exc: ready_layer = None local_error = exc - ready_count = self._control_count( + ready_count = self._control_ready_count.fill_( EPLB_CONTROL_ERROR if local_error is not None else int(ready_layer is not None) ) dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) @@ -341,7 +313,6 @@ def _finish_rebalance(self): self.current_placement = self.target_placement self.target_placement = None self.target_metadata = None - self.in_flight = False self._restart_collection() if self.global_rank == 0: logger.info( diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index faa6ca46b4..afba43a7ad 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1,5 +1,6 @@ import threading import time +from concurrent.futures import Future from contextlib import nullcontext from types import SimpleNamespace @@ -666,9 +667,11 @@ def test_steady_sampling_resets_aggregated_route_counter(): world_size=1, ) manager._eplb_impls = [impl] + manager.prefill_steps = 0 + manager.step_interval = 20 - manager._reset_route_counters() - manager._reset_route_counters() + manager._restart_collection() + manager._restart_collection() assert impl.route_counter.shape == (4,) assert torch.count_nonzero(impl.route_counter) == 0 @@ -721,9 +724,6 @@ def test_manager_evaluation_collective_sums_rank_loads(monkeypatch): manager.num_redundant_experts_per_rank = 1 manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0) manager.evaluation_group = object() - manager._evaluation_lock = threading.Lock() - manager._evaluation_result = None - manager._evaluation_error = None local = torch.full((1, 4), 100, dtype=torch.int64) manager._collect_local_load = lambda: local.clone() seen = {} @@ -743,14 +743,15 @@ def plan_and_broadcast(global_load): manager._plan_and_broadcast = plan_and_broadcast - manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) + evaluation = Future() + manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})(), evaluation) + result = evaluation.result() assert seen["group"] is manager.evaluation_group assert seen["before"].shape == (1, 4) assert torch.equal(seen["before"], local) assert torch.equal(seen["global_load"], torch.full_like(local, 200)) - assert manager._evaluation_error is None - assert manager._evaluation_result["expert_imbalance_ratio"] == 1.0 + assert result["expert_imbalance_ratio"] == 1.0 def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_call( @@ -781,9 +782,6 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c manager.num_redundant_experts_per_rank = 1 manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() manager.evaluation_group = object() - manager._evaluation_lock = threading.Lock() - manager._evaluation_result = None - manager._evaluation_error = None manager._collect_local_load = lambda: torch.full((3, 4), 100, dtype=torch.int64) planned_placement = torch.tensor( [ @@ -814,13 +812,14 @@ def build_maps_for_layers(*args, **kwargs): build_maps_for_layers, ) - manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) + evaluation = Future() + manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})(), evaluation) + result = evaluation.result() - assert manager._evaluation_error is None assert calls == [(2, 4, 2)] - metadata = manager._evaluation_result["metadata"] - assert metadata[1] is None - assert [layer_index for layer_index, _plan in manager._evaluation_result["layer_plans"]] == [0, 2] + metadata = result["metadata"] + assert 1 not in metadata + assert [layer_index for layer_index, _plan in result["layer_plans"]] == [0, 2] for layer_index in (0, 2): item = metadata[layer_index] expected = build_logical_to_physical_map( @@ -1407,11 +1406,10 @@ def remote_error(tensor, **_kwargs): def test_evaluation_ready_gate_propagates_local_and_remote_errors(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager._evaluation_lock = threading.Lock() manager.control_group = object() manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager._evaluation_error = RuntimeError("evaluation boom") - manager._evaluation_result = None + manager._evaluation = Future() + manager._evaluation.set_exception(RuntimeError("evaluation boom")) statuses = [] def retain_local_error(tensor, **_kwargs): @@ -1424,8 +1422,8 @@ def retain_local_error(tensor, **_kwargs): assert str(exc_info.value.__cause__) == "evaluation boom" assert statuses == [manager_module.EPLB_CONTROL_ERROR] - manager._evaluation_error = None - manager._evaluation_result = {"kind": "no_improvement"} + manager._evaluation = Future() + manager._evaluation.set_result({"kind": "no_improvement"}) statuses.clear() def remote_error(tensor, **_kwargs): @@ -1500,7 +1498,8 @@ def finish(self): def test_manager_inflight_step_does_not_poll(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = True + manager.in_flight_layers = [0] + manager._evaluation = None calls = [] manager._poll_in_flight = lambda: calls.append("poll") manager.step() @@ -1509,8 +1508,8 @@ def test_manager_inflight_step_does_not_poll(): def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.evaluation_in_flight = False + manager.in_flight_layers = [] + manager._evaluation = None manager.prefill_steps = 0 manager.step_interval = 3 manager.next_evaluation_step = 3 @@ -1548,8 +1547,8 @@ def test_manager_restart_collection_resets_counters_and_arms_next_window(): def test_manager_poll_advances_inflight_before_evaluation(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = True - manager.evaluation_in_flight = True + manager.in_flight_layers = [0] + manager._evaluation = Future() calls = [] manager._poll_in_flight = lambda: calls.append("inflight") manager._finish_evaluation = lambda: calls.append("evaluation") @@ -1561,11 +1560,8 @@ def test_manager_poll_advances_inflight_before_evaluation(): def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_result = None - manager._evaluation_error = None + manager.in_flight_layers = [] + manager._evaluation = Future() manager.control_group = object() manager._control_ready_count = torch.empty(1, dtype=torch.int32) calls = [] @@ -1575,7 +1571,7 @@ def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): manager.poll() assert calls == [] - manager._evaluation_result = {"kind": "no_improvement"} + manager._evaluation.set_result({"kind": "no_improvement"}) monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) manager.poll() assert calls == ["evaluation"] From 4c5046df6f63177a31f00202dc4e0fe85936f11d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 06:53:59 +0000 Subject: [PATCH 27/72] refactor(eplb): simplify expert p2p transfer --- lightllm/common/eplb_utils.py | 38 +- .../model_infer/mode_backend/eplb_manager.py | 225 ++++---- .../model_infer/mode_backend/eplb_transfer.py | 437 +++++++------- unit_tests/common/fused_moe/test_eplb.py | 531 +++++++++--------- .../fused_moe/test_eplb_transfer_gpu.py | 66 ++- 5 files changed, 692 insertions(+), 605 deletions(-) diff --git a/lightllm/common/eplb_utils.py b/lightllm/common/eplb_utils.py index 4dd901a600..073847292a 100644 --- a/lightllm/common/eplb_utils.py +++ b/lightllm/common/eplb_utils.py @@ -1,13 +1,35 @@ -"""Small, dependency-free EPLB helpers shared by transfer and model profiling.""" +"""传输与模型分析共用的轻量 EPLB 工具。""" +from typing import List, Optional, Protocol, Tuple -def extract_eplb_expert_tensors(weight): - result = [] +import torch + +NamedTensor = Tuple[str, torch.Tensor] + + +class ExpertWeightPack(Protocol): + """EPLB 需要迁移的单组专家权重。""" + + weight: torch.Tensor + weight_scale: Optional[torch.Tensor] + weight_zero_point: Optional[torch.Tensor] + + +class EPLBExpertWeight(Protocol): + """包含门控投影和下投影专家权重的 MoE 层。""" + + w13: ExpertWeightPack + w2: ExpertWeightPack + + +def extract_eplb_expert_tensors(weight: EPLBExpertWeight) -> List[NamedTensor]: + """按固定顺序返回 EPLB 必须迁移的权重及量化参数。""" + named_tensors: List[NamedTensor] = [] for pack_name in ("w13", "w2"): - pack = getattr(weight, pack_name) - for value_name in ("weight", "weight_scale"): - tensor = getattr(pack, value_name, None) + weight_pack: ExpertWeightPack = getattr(weight, pack_name) + for value_name in ("weight", "weight_scale", "weight_zero_point"): + tensor: Optional[torch.Tensor] = getattr(weight_pack, value_name, None) if tensor is not None: assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" - result.append((f"{pack_name}.{value_name}", tensor)) - return result + named_tensors.append((f"{pack_name}.{value_name}", tensor)) + return named_tensors diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 6601fe7e84..09701f7b83 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -1,6 +1,7 @@ from concurrent.futures import Future import threading import time +from typing import Any, Dict, List, Optional, Tuple import torch import torch.distributed as dist @@ -19,6 +20,7 @@ ) from lightllm.server.metrics.manager import MetricClient from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + EPLBTransferInfo, PinnedMemoryEPLBTransfer, build_transfer_plan, ) @@ -41,35 +43,35 @@ class EPLBManager: - """Collect logical load, invoke a planner, and publish migrated rows.""" + """收集专家负载、调用布局规划器,并在安全边界发布迁移后的专家权重。""" - def __init__(self, model: TpPartBaseModel): - weights = _find_fused_moe_weights(model) + def __init__(self, model: TpPartBaseModel) -> None: + weights: List[FusedMoeWeight] = _find_fused_moe_weights(model) assert weights, "EPLB requires at least one EP MoE layer" - self.global_rank = get_global_rank() - self.world_size = get_global_world_size() - self.node_world_size = get_node_world_size() + self._weights: List[FusedMoeWeight] = weights + self.global_rank: int = get_global_rank() + self.world_size: int = get_global_world_size() + self.node_world_size: int = get_node_world_size() self._eplb_impls = [weight.fuse_moe_impl for weight in weights] - routed = {impl.n_routed_experts for impl in self._eplb_impls} redundant = {impl.num_redundant_experts_per_rank for impl in self._eplb_impls} assert len(routed) == len(redundant) == 1 - self.num_logical_experts = routed.pop() - self.num_redundant_experts_per_rank = redundant.pop() - self.step_interval = get_prefill_eplb_step_interval() - self.prefill_steps = 0 - self.next_evaluation_step = self.step_interval + self.num_logical_experts: int = routed.pop() + self.num_redundant_experts_per_rank: int = redundant.pop() + self.num_primary_experts_per_rank: int = self.num_logical_experts // self.world_size + self.step_interval: int = get_prefill_eplb_step_interval() + self.prefill_steps: int = 0 + self.next_evaluation_step: int = self.step_interval initial_local_expert_ids = build_initial_local_expert_ids( self.num_logical_experts, self.world_size, self.num_redundant_experts_per_rank, ) - num_primary_experts_per_rank = self.num_logical_experts // self.world_size initial_redundant_expert_ids = [ - expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids + expert_ids[self.num_primary_experts_per_rank :] for expert_ids in initial_local_expert_ids ] - self.current_placement = torch.tensor( + self.current_placement: torch.Tensor = torch.tensor( [initial_redundant_expert_ids for _ in weights], dtype=torch.int64, ) @@ -80,17 +82,18 @@ def __init__(self, model: TpPartBaseModel): rebalance_gain_threshold=get_eplb_rebalance_gain_threshold(), ) - self.in_flight_layers = [] - self.target_placement = None - self.target_metadata = None - self._evaluation = None - self.metric_client = None + self.in_flight_transfers: List[EPLBTransferInfo] = [] + self.completed_transfers: List[PinnedMemoryEPLBTransfer] = [] + self.active_transfer: Optional[PinnedMemoryEPLBTransfer] = None + self.target_placement: Optional[torch.Tensor] = None + self.target_metadata: Optional[Dict[int, torch.Tensor]] = None + self._evaluation: Optional[Future] = None + self.metric_client: Optional[MetricClient] = None self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") - self._control_ready_count = torch.empty(1, dtype=torch.int32) - self.transfer = PinnedMemoryEPLBTransfer(weights, self.transfer_group, self.global_rank) + self._control_status_buffer: torch.Tensor = torch.empty(1, dtype=torch.int32) self._restart_collection() if self.global_rank == 0: @@ -100,22 +103,22 @@ def __init__(self, model: TpPartBaseModel): f"step_interval={self.step_interval} planner={type(self.planner).__name__}" ) - def poll(self): - """Advance EPLB only from a globally ordered pre-forward boundary.""" - if self.in_flight_layers: + def poll(self) -> None: + """在所有 rank 顺序一致的推理边界推进评估或专家传输。""" + if self.in_flight_transfers: self._poll_in_flight() elif self._evaluation is not None and self._evaluation_ready_on_all_ranks(): self._finish_evaluation() - def step(self): - if self.in_flight_layers or self._evaluation is not None: + def step(self) -> None: + if self.in_flight_transfers or self._evaluation is not None: return self.prefill_steps += 1 if self.prefill_steps >= self.next_evaluation_step: self._start_evaluation() - def _restart_collection(self): - counters = [impl.route_counter for impl in self._eplb_impls] + def _restart_collection(self) -> None: + counters: List[torch.Tensor] = [impl.route_counter for impl in self._eplb_impls] torch._foreach_zero_(counters) for impl in self._eplb_impls: impl.recording = True @@ -127,9 +130,9 @@ def _collect_local_load(self) -> torch.Tensor: raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") return torch.stack(counters).cpu() - def _plan_and_broadcast(self, global_load: torch.Tensor): - result = None - local_error = None + def _plan_and_broadcast(self, global_load: torch.Tensor) -> Dict[str, Any]: + result: Optional[Dict[str, Any]] = None + local_error: Optional[BaseException] = None if self.global_rank == 0: try: result = self.planner.plan( @@ -149,15 +152,17 @@ def _plan_and_broadcast(self, global_load: torch.Tensor): raise RuntimeError(f"EPLB planner failed on rank zero: {result['message']}") return result - def _build_rebalance_data(self, result): - metadata = {} - layer_plans = [] - changed_layer_indices = [layer_index for layer_index, changed in enumerate(result["changed_layers"]) if changed] + def _build_rebalance_data(self, result: Dict[str, Any]) -> Tuple[Dict[int, torch.Tensor], List[EPLBTransferInfo]]: + metadata_by_layer: Dict[int, torch.Tensor] = {} + planned_transfers: List[EPLBTransferInfo] = [] + changed_layer_indices: List[int] = [ + layer_index for layer_index, changed in enumerate(result["changed_layers"]) if changed + ] num_primary_experts_per_rank = self.num_logical_experts // self.world_size - rank_to_logic_expert_ids_by_layer = [] + local_expert_ids_by_rank_and_layer: List[List[List[int]]] = [] for layer_index in changed_layer_indices: - rank_to_logic_expert_ids_by_layer.append( + local_expert_ids_by_rank_and_layer.append( [ list( range( @@ -169,32 +174,30 @@ def _build_rebalance_data(self, result): for rank in range(self.world_size) ] ) - maps = torch.tensor( + logical_to_physical_maps = torch.tensor( build_logical_to_physical_maps_for_layers( - rank_to_logic_expert_ids_by_layer, + local_expert_ids_by_rank_and_layer, self.num_logical_experts, current_rank=self.global_rank, ), dtype=torch.int32, ) - for offset, layer_index in enumerate(changed_layer_indices): - target = torch.tensor(result["placement"][layer_index], dtype=torch.int64) - metadata[layer_index] = maps[offset] - layer_plans.append( - ( + for changed_layer_offset, layer_index in enumerate(changed_layer_indices): + target_layer_placement = torch.tensor(result["placement"][layer_index], dtype=torch.int64) + metadata_by_layer[layer_index] = logical_to_physical_maps[changed_layer_offset] + planned_transfers.extend( + build_transfer_plan( + self.current_placement[layer_index], + target_layer_placement, layer_index, - build_transfer_plan( - self.current_placement[layer_index], - target, - self.num_logical_experts, - self.world_size, - self.node_world_size, - ), + self.num_logical_experts, + self.world_size, + self.node_world_size, ) ) - return metadata, layer_plans + return metadata_by_layer, planned_transfers - def _evaluate_after_event(self, event: torch.cuda.Event, evaluation: Future): + def _evaluate_after_event(self, event: torch.cuda.Event, evaluation: Future) -> None: try: torch.cuda.set_device(self._eplb_impls[0].route_counter.device) event.synchronize() @@ -203,12 +206,12 @@ def _evaluate_after_event(self, event: torch.cuda.Event, evaluation: Future): result = self._plan_and_broadcast(global_load) result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) if result["kind"] == "planned": - result["metadata"], result["layer_plans"] = self._build_rebalance_data(result) + result["metadata"], result["transfer_infos"] = self._build_rebalance_data(result) evaluation.set_result(result) except BaseException as exc: evaluation.set_exception(exc) - def _start_evaluation(self): + def _start_evaluation(self) -> None: for impl in self._eplb_impls: impl.recording = False event = torch.cuda.Event() @@ -224,16 +227,16 @@ def _evaluation_ready_on_all_ranks(self) -> bool: ready = self._evaluation.done() error = self._evaluation.exception() if ready else None status = EPLB_CONTROL_ERROR if error is not None else int(ready) - ready_count = self._control_ready_count.fill_(status) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) - ready_count = int(ready_count.item()) - if ready_count < 0: + global_evaluation_status = self._control_status_buffer.fill_(status) + dist.all_reduce(global_evaluation_status, op=dist.ReduceOp.MIN, group=self.control_group) + global_evaluation_status_value = int(global_evaluation_status.item()) + if global_evaluation_status_value < 0: if error is not None: raise RuntimeError("EPLB evaluation failed on this rank") from error raise RuntimeError("EPLB evaluation failed on another rank") - return bool(ready_count) + return bool(global_evaluation_status_value) - def _finish_evaluation(self): + def _finish_evaluation(self) -> None: result = self._evaluation.result() self._evaluation = None self._publish_expert_load_metric(result) @@ -244,21 +247,22 @@ def _finish_evaluation(self): return self._start_rebalance(result) - def _publish_expert_load_metric(self, result): + def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: if self.global_rank != 0: return if self.metric_client is None: self.metric_client = MetricClient(get_shm_port_args().metric_port) self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) - def _start_rebalance(self, result): + def _start_rebalance(self, result: Dict[str, Any]) -> None: self.target_placement = torch.tensor(result["placement"], dtype=torch.int64) self.target_metadata = result["metadata"] - self.in_flight_layers = [layer_index for layer_index, _ in result["layer_plans"]] + self.in_flight_transfers = result["transfer_infos"] + self.completed_transfers = [] self.in_flight_started_at = time.time() - self.transfer.start(result["layer_plans"]) + self._start_next_transfer() if self.global_rank == 0: - changed_slots = sum(len(plan) for _, plan in result["layer_plans"]) + changed_slots = len(self.in_flight_transfers) logger.info( "eplb started prefill_steps=%s max_before=%.4f max_after=%.4f " "p95_before=%.4f p95_after=%.4f rebalance_gain=%.4f " @@ -273,43 +277,68 @@ def _start_rebalance(self, result): changed_slots, ) - def _poll_in_flight(self): - local_error = None - try: - ready_layer = self.transfer.ready_layer() - except BaseException as exc: - ready_layer = None - local_error = exc - ready_count = self._control_ready_count.fill_( - EPLB_CONTROL_ERROR if local_error is not None else int(ready_layer is not None) - ) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) - ready_count = int(ready_count.item()) - if ready_count < 0: - if local_error is not None: - raise RuntimeError("EPLB transfer worker failed on this rank") from local_error - raise RuntimeError("EPLB transfer worker failed on another rank") - if ready_count == 0: + def _poll_in_flight(self) -> None: + assert self.active_transfer is not None + all_ranks_finished = self._control_status_buffer.fill_(int(self.active_transfer.is_finished())) + dist.all_reduce(all_ranks_finished, op=dist.ReduceOp.MIN, group=self.control_group) + if not bool(all_ranks_finished.item()): + return + expected_transfer_info: EPLBTransferInfo = self.in_flight_transfers[0] + if self.active_transfer.transfer_info != expected_transfer_info: + raise RuntimeError("EPLB completed transfer does not match the expected transfer info") + + self.completed_transfers.append(self.active_transfer) + self.in_flight_transfers.pop(0) + + if self.in_flight_transfers and self.in_flight_transfers[0].layer_index == expected_transfer_info.layer_index: + self._start_next_transfer() return - if ready_layer != self.in_flight_layers[0]: - raise RuntimeError(f"EPLB ready layer {ready_layer} does not match expected {self.in_flight_layers[0]}") from lightllm.server.router.model_infer.infer_batch import g_infer_context torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - self.transfer.commit(ready_layer, lambda: self._commit_layer_metadata(ready_layer)) - self.in_flight_layers.pop(0) - if not self.in_flight_layers: - self.transfer.finish() + layer_index: int = expected_transfer_info.layer_index + self._commit_transferred_layer(layer_index) + self.completed_transfers.clear() + if self.in_flight_transfers: + self._start_next_transfer() + else: + self.active_transfer = None self._finish_rebalance() - def _commit_layer_metadata(self, layer_index: int): - self._eplb_impls[layer_index].logical_to_physical_map.copy_( - self.target_metadata[layer_index], - non_blocking=True, + def _start_next_transfer(self) -> None: + transfer_info: EPLBTransferInfo = self.in_flight_transfers[0] + self.active_transfer = PinnedMemoryEPLBTransfer( + self._weights, + self.transfer_group, + self.global_rank, + transfer_info, ) - - def _finish_rebalance(self): + self.active_transfer.start() + + def _commit_transferred_layer(self, layer_index: int) -> None: + """在主推理线程中同步发布一层权重和路由 metadata。""" + assert self.target_placement is not None + target_redundant_expert_ids: List[int] = self.target_placement[layer_index, self.global_rank].tolist() + for transfer in self.completed_transfers: + transfer_info: EPLBTransferInfo = transfer.transfer_info + if transfer_info.dest_rank != self.global_rank: + continue + destination_slot_index: int = target_redundant_expert_ids.index(transfer_info.source_logical_expert_id) + destination_local_expert_index: int = self.num_primary_experts_per_rank + destination_slot_index + for tensor_buffer in transfer.tensor_buffers: + tensor_buffer.live_tensor[destination_local_expert_index].copy_(tensor_buffer.pinned_row) + + local_expert_ids: List[int] = self._eplb_impls[layer_index].local_logics_expert_ids_list + local_expert_ids[self.num_primary_experts_per_rank :] = target_redundant_expert_ids + self._commit_layer_metadata(layer_index) + + def _commit_layer_metadata(self, layer_index: int) -> None: + assert self.target_metadata is not None + self._eplb_impls[layer_index].logical_to_physical_map.copy_(self.target_metadata[layer_index]) + + def _finish_rebalance(self) -> None: + assert self.target_placement is not None self.current_placement = self.target_placement self.target_placement = None self.target_metadata = None @@ -334,8 +363,8 @@ def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: return float(ratios.mean().item()) -def _find_fused_moe_weights(model): - weights_by_id = {} +def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: + weights_by_id: Dict[int, FusedMoeWeight] = {} for layer in model.trans_layers_weight: for value in getattr(layer, "__dict__", {}).values(): if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index 32477f1687..05f0cbf0ff 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -1,225 +1,270 @@ -"""Layer-by-layer expert-row migration for EPLB.""" +"""EPLB 专家权重的逐层迁移。""" +import os import threading -from collections import defaultdict +import zlib from dataclasses import dataclass -from typing import List, Sequence, Tuple +from enum import Enum +from typing import List, Sequence import torch import torch.distributed as dist -from lightllm.common.eplb_utils import extract_eplb_expert_tensors +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight +from lightllm.common.eplb_utils import NamedTensor, extract_eplb_expert_tensors +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +@dataclass(frozen=True) +class ExpertTensorBuffer: + """一项 live 专家张量及其单专家 pinned memory 缓冲行。""" + + name: str + live_tensor: torch.Tensor + pinned_row: torch.Tensor + + +@dataclass(frozen=True) +class EPLBTransferInfo: + """单个逻辑专家的一次传输描述。 + + ``layer_index`` 是专家权重在 EPLB 层列表中的下标。源 rank 使用 + ``source_logical_expert_id`` 定位当前本地物理行,目标 rank 从 + ``tensor_buffers`` 中读取传输完成的 pinned memory 数据。 + """ + + source_rank: int + layer_index: int + source_logical_expert_id: int + dest_rank: int @dataclass(frozen=True) -class TransferStep: - dst_rank: int - dst_slot: int - src_rank: int - src_local_row: int +class _ExpertSource: + """一个逻辑专家当前可用的物理副本位置。""" + + rank: int + local_expert_index: int + + +class TransferStatus(Enum): + """异步传输线程的生命周期状态。""" + + IDLE = "idle" + RUNNING = "running" + SUCCEEDED = "succeeded" def build_transfer_plan( - current: torch.Tensor, - target: torch.Tensor, + current_placement: torch.Tensor, + target_placement: torch.Tensor, + layer_index: int, num_logical_experts: int, world_size: int, node_world_size: int, -) -> List[TransferStep]: - assert tuple(current.shape) == tuple(target.shape) == (world_size, current.shape[1]) - num_experts_per_rank = num_logical_experts // world_size - current_rows = current.tolist() - target_rows = target.tolist() - # A primary row is always a valid source. Existing replicas are also - # candidates so that a destination can prefer a same-node copy. - candidates_by_expert = [ - [(expert // num_experts_per_rank, expert % num_experts_per_rank)] for expert in range(num_logical_experts) - ] - for rank, row in enumerate(current_rows): - for slot, expert in enumerate(row): - candidates_by_expert[expert].append((rank, num_experts_per_rank + slot)) - source_load = [0] * world_size - plan = [] - for dst_rank in range(world_size): - for dst_slot, expert in enumerate(target_rows[dst_rank]): - if expert == current_rows[dst_rank][dst_slot]: +) -> List[EPLBTransferInfo]: + """根据新旧冗余专家分布生成确定性的传输计划。 + + ``current_placement`` 和 ``target_placement`` 只描述冗余槽位,形状均为 + ``[world_size, num_redundant_slots]``。固定主专家不在这两个张量中,但始终 + 可以作为数据源。返回结果只包含发生变化的目标槽位所需专家,每个 + :class:`EPLBTransferInfo` 只描述一个逻辑专家的传输。 + + 为同一个逻辑专家选择数据源时,依次考虑: + + 1. 优先使用目标 rank 所在节点上的已有副本,避免跨节点传输; + 2. 均衡各个源 rank 承担的传输次数; + 3. 使用 rank 和本地行号做稳定排序,保证所有 rank 生成一致结果。 + """ + assert ( + tuple(current_placement.shape) + == tuple(target_placement.shape) + == ( + world_size, + current_placement.shape[1], + ) + ) + num_primary_experts_per_rank = num_logical_experts // world_size + current_placement_by_rank: List[List[int]] = current_placement.tolist() + target_placement_by_rank: List[List[int]] = target_placement.tolist() + + # 每个逻辑专家的主专家行永远存在,因此先把它加入候选源;当前仍存在的 + # 冗余副本也可以作为源,这样目标 rank 有机会直接使用同节点副本。 + source_candidates_by_expert: List[List[_ExpertSource]] = [] + for logical_expert_id in range(num_logical_experts): + primary_rank, primary_local_expert_index = divmod(logical_expert_id, num_primary_experts_per_rank) + source_candidates_by_expert.append([_ExpertSource(primary_rank, primary_local_expert_index)]) + for rank, redundant_expert_ids in enumerate(current_placement_by_rank): + for redundant_slot_index, logical_expert_id in enumerate(redundant_expert_ids): + source_candidates_by_expert[logical_expert_id].append( + _ExpertSource( + rank=rank, + local_expert_index=num_primary_experts_per_rank + redundant_slot_index, + ) + ) + + num_transfers_by_source_rank: List[int] = [0] * world_size + transfer_infos: List[EPLBTransferInfo] = [] + for destination_rank in range(world_size): + for destination_slot_index, logical_expert_id in enumerate(target_placement_by_rank[destination_rank]): + if logical_expert_id == current_placement_by_rank[destination_rank][destination_slot_index]: continue - src_rank, src_row = min( - candidates_by_expert[expert], - key=lambda item: ( - item[0] // node_world_size != dst_rank // node_world_size, - source_load[item[0]], - item[0], - item[1], + source = min( + source_candidates_by_expert[logical_expert_id], + key=lambda candidate: ( + candidate.rank // node_world_size != destination_rank // node_world_size, + num_transfers_by_source_rank[candidate.rank], + candidate.rank, + candidate.local_expert_index, ), ) - source_load[src_rank] += 1 - plan.append(TransferStep(dst_rank, dst_slot, src_rank, src_row)) - return plan + num_transfers_by_source_rank[source.rank] += 1 + transfer_infos.append( + EPLBTransferInfo( + source_rank=source.rank, + layer_index=layer_index, + source_logical_expert_id=logical_expert_id, + dest_rank=destination_rank, + ) + ) + + return transfer_infos class PinnedMemoryEPLBTransfer: - """Move one layer at a time through reusable pinned CPU row buffers. + """在后台线程中传输一个逻辑专家的全部权重张量。 + + 只有源 rank 和目标 rank 参与 Gloo 点对点通信,具体的数据路径是: + + ``源 GPU 权重行 -> 源 rank 的 pinned CPU 行 -> 目标 rank 的 pinned CPU 行`` + + 如果源和目标是同一个 rank,则只执行 GPU 到 pinned CPU 的本地复制,不产生 + 网络通信。其他 rank 不分配 pinned row,也不参与该专家的数据传输。 - Every rank executes the same ordered Gloo broadcasts. A source first - copies a live GPU row to pinned memory; destinations then copy that row - into a single-layer GPU staging buffer. The inference thread publishes - the staging rows and routing metadata together at a safe forward boundary. + 本类只负责异步传输,不修改 live 权重,也不更新路由 metadata。传输成功后, + ``status`` 会变为 :attr:`TransferStatus.SUCCEEDED`,收到的数据保存在 + ``tensor_buffers``。EPLBManager 在主循环的安全边界同步提交这些数据。 + + 每个对象只表示构造函数中 ``transfer_info`` 指定的一次传输。逻辑专家 ID + 在源 rank 上通过该层当前的本地专家列表解析为物理行;目标物理槽位不属于 + 传输职责,由 manager 根据目标 placement 决定。 """ - def __init__(self, weights, transfer_group, global_rank): - self._eplb_impls = [weight.fuse_moe_impl for weight in weights] - self.transfer_group = transfer_group - self.global_rank = global_rank - self.num_experts_per_rank = self._eplb_impls[0].num_primary_experts_per_rank - self.device = weights[0].w13.weight.device - self.live = [extract_eplb_expert_tensors(weight) for weight in weights] - self._validate_live_layout() - - num_redundant_slots = self._eplb_impls[0].num_redundant_experts_per_rank - self.staging = [ - ( - name, - torch.empty( - (num_redundant_slots,) + tuple(tensor.shape[1:]), - dtype=tensor.dtype, - device=tensor.device, - ), - ) - for name, tensor in self.live[0] - ] - self.pinned_rows = [ - ( - name, - torch.empty( - tuple(tensor.shape[1:]), - dtype=tensor.dtype, + def __init__( + self, + weights: Sequence[FusedMoeWeight], + transfer_group: dist.ProcessGroup, + current_global_rank: int, + transfer_info: EPLBTransferInfo, + ) -> None: + self._p2p_group: dist.ProcessGroup = transfer_group + self._is_source_rank: bool = current_global_rank == transfer_info.source_rank + self._is_destination_rank: bool = current_global_rank == transfer_info.dest_rank + self.transfer_info: EPLBTransferInfo = transfer_info + + layer_weight: FusedMoeWeight = weights[transfer_info.layer_index] + # 提取该 MoE 层实际参与推理的专家张量,包括 w13/w2 的量化后权重 + # (或非量化权重),以及配套的 weight_scale、weight_zero_point 等量化 + # 信息。后续会为每项张量创建对应的 pinned row,确保专家状态完整迁移。 + named_live_tensors: List[NamedTensor] = extract_eplb_expert_tensors(layer_weight) + self._local_logical_expert_ids: List[int] = layer_weight.fuse_moe_impl.local_logics_expert_ids_list + self._device: torch.device = named_live_tensors[0][1].device + + # 只有源和目标 rank 需要保存该专家的 pinned row。源 rank 用它作为 + # send buffer,目标 rank 用它作为 recv buffer,并在传输完成后直接交给 + # manager 提交到 live 权重,避免无关 rank 分配同样大小的 pinned memory。 + self.tensor_buffers: List[ExpertTensorBuffer] = [] + if self._is_source_rank or self._is_destination_rank: + for tensor_name, live_tensor in named_live_tensors: + pinned_row = torch.empty( + tuple(live_tensor.shape[1:]), + dtype=live_tensor.dtype, device="cpu", pin_memory=True, - ), - ) - for name, tensor in self.live[0] - ] - self._copy_stream = torch.cuda.Stream(device=self.device) - self._release = threading.Event() - self._release.set() - self._consumed_event = torch.cuda.Event() - self._consumed_recorded = False - self._ready = None - self._ready_lock = threading.Lock() - self._error = None - self._thread = None - - def _validate_live_layout(self) -> None: - reference = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in self.live[0]] - num_redundant_slots = self._eplb_impls[0].num_redundant_experts_per_rank - for layer_index, (impl, tensors) in enumerate(zip(self._eplb_impls, self.live)): - layout = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in tensors] - assert layout == reference, f"EPLB layer {layer_index} has incompatible expert tensor layout" - assert impl.num_redundant_experts_per_rank == num_redundant_slots, "EPLB redundant slot count must match" - - @staticmethod - def _group_steps_by_source(plan: Sequence[TransferStep]): - grouped = defaultdict(list) - for step in plan: - grouped[(step.src_rank, step.src_local_row)].append(step) - return [(source, grouped[source]) for source in sorted(grouped)] - - def _copy_layer(self, layer_index: int, plan: Sequence[TransferStep]) -> None: - live_tensors = self.live[layer_index] - for (src_rank, src_local_row), steps in self._group_steps_by_source(plan): - if self.global_rank == src_rank: - with torch.cuda.stream(self._copy_stream): - for (_, live), (_, pinned) in zip(live_tensors, self.pinned_rows): - pinned.copy_(live[src_local_row], non_blocking=True) - # Gloo must not read the CPU row before the device-to-host copy completes. - self._copy_stream.synchronize() - for _, pinned in self.pinned_rows: - dist.broadcast(pinned, src=src_rank, group=self.transfer_group) - - dst_slots = sorted({step.dst_slot for step in steps if step.dst_rank == self.global_rank}) - if dst_slots: - with torch.cuda.stream(self._copy_stream): - for (_, staging), (_, pinned) in zip(self.staging, self.pinned_rows): - for dst_slot in dst_slots: - staging[dst_slot].copy_(pinned, non_blocking=True) - self._copy_stream.synchronize() - - def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]) -> None: - if self._thread is not None: - raise RuntimeError("EPLB transfer has not been finished") - self._error = None - with self._ready_lock: - if self._ready is not None: - raise RuntimeError("EPLB ready layer has not been committed") - - def worker() -> None: - try: - torch.cuda.set_device(self.device) - for layer_index, plan in layer_plans: - self._release.wait() - self._release.clear() - if self._consumed_recorded: - self._consumed_event.synchronize() - changed_dst_slots = tuple( - sorted({step.dst_slot for step in plan if step.dst_rank == self.global_rank}) + ) + self.tensor_buffers.append( + ExpertTensorBuffer( + name=tensor_name, + live_tensor=live_tensor, + pinned_row=pinned_row, ) - self._copy_layer(layer_index, plan) - with self._ready_lock: - self._ready = (layer_index, changed_dst_slots) - except BaseException as exc: - self._error = exc - - self._thread = threading.Thread(target=worker, name="eplb-pin-memory", daemon=True) - self._thread.start() - - def ready_layer(self): - if self._error is not None: - raise RuntimeError("EPLB migration worker failed") from self._error - with self._ready_lock: - return None if self._ready is None else self._ready[0] - - def commit(self, layer_index: int, post_copy=None) -> None: - with self._ready_lock: - if self._ready is None or self._ready[0] != layer_index: - raise RuntimeError("EPLB commit does not match the ready layer") - _, changed_dst_slots = self._ready - self._ready = None - for (_, live), (_, staging) in zip(self.live[layer_index], self.staging): - _commit_staging_rows(live, staging, self.num_experts_per_rank, changed_dst_slots) - if post_copy is not None: - post_copy() - self._consumed_event.record(torch.cuda.current_stream()) - self._consumed_recorded = True - self._release.set() - - def finish(self) -> None: - thread = self._thread - if thread is None: - return - thread.join() - self._thread = None - if self._error is not None: - raise RuntimeError("EPLB migration worker failed") from self._error - - -def _commit_staging_rows( - live: torch.Tensor, - staging: torch.Tensor, - num_experts_per_rank: int, - changed_dst_slots: Sequence[int], -) -> None: - slots = sorted(set(changed_dst_slots)) - if not slots: - return - run_start = previous = slots[0] - for dst_slot in (*slots[1:], None): - if dst_slot is not None and dst_slot == previous + 1: - previous = dst_slot - continue - run_length = previous - run_start + 1 - live.narrow(0, num_experts_per_rank + run_start, run_length).copy_( - staging.narrow(0, run_start, run_length), non_blocking=True + ) + self._device_to_host_stream: torch.cuda.Stream = torch.cuda.Stream(device=self._device) + + self.status: TransferStatus = TransferStatus.IDLE + self._transfer_thread: threading.Thread = threading.Thread( + target=self._run_transfer, + name=f"eplb-transfer-layer-{transfer_info.layer_index}-expert-{transfer_info.source_logical_expert_id}", + daemon=True, + ) + + def start(self) -> None: + """启动构造函数中 transfer_info 描述的异步传输。""" + assert self.status is TransferStatus.IDLE, "EPLB transfer has already been started" + self.status = TransferStatus.RUNNING + self._transfer_thread.start() + + def is_finished(self) -> bool: + """返回后台传输是否已经成功完成。""" + return self.status is TransferStatus.SUCCEEDED + + def _run_transfer(self) -> None: + """把指定专家的全部权重行传输到各 rank 的 pinned memory。""" + try: + transfer_info: EPLBTransferInfo = self.transfer_info + if self._is_source_rank: + torch.cuda.set_device(self._device) + source_local_expert_index: int = self._local_logical_expert_ids.index( + transfer_info.source_logical_expert_id + ) + with torch.cuda.stream(self._device_to_host_stream): + for tensor_buffer in self.tensor_buffers: + tensor_buffer.pinned_row.copy_( + tensor_buffer.live_tensor[source_local_expert_index], + non_blocking=True, + ) + # Gloo 读取 pinned row 前,源 rank 必须等待 GPU -> CPU 拷贝完成。 + self._device_to_host_stream.synchronize() + + if transfer_info.source_rank != transfer_info.dest_rank: + if self._is_source_rank: + for tensor_buffer in self.tensor_buffers: + message_tag = self._build_p2p_message_tag(tensor_buffer.name) + dist.send( + tensor_buffer.pinned_row, + dst=transfer_info.dest_rank, + group=self._p2p_group, + tag=message_tag, + ) + elif self._is_destination_rank: + for tensor_buffer in self.tensor_buffers: + message_tag = self._build_p2p_message_tag(tensor_buffer.name) + dist.recv( + tensor_buffer.pinned_row, + src=transfer_info.source_rank, + group=self._p2p_group, + tag=message_tag, + ) + self.status = TransferStatus.SUCCEEDED + except BaseException: + logger.exception("EPLB transfer failed") + os._exit(1) + + def _build_p2p_message_tag(self, tensor_name: str) -> int: + """为当前专家张量生成 source 和 destination 一致的 Gloo 整数 tag。 + + Python ``hash`` 会因进程随机种子不同而产生不同结果,因此这里使用稳定的 + CRC32,并限制到 Gloo 可安全使用的有符号 31 位整数范围。标识中包含层、 + 源 rank、目标 rank、逻辑专家和张量名称,避免依赖张量列表的隐式顺序。 + """ + transfer_info: EPLBTransferInfo = self.transfer_info + message_identity = ( + f"{transfer_info.layer_index}:" + f"{transfer_info.source_rank}:" + f"{transfer_info.dest_rank}:" + f"{transfer_info.source_logical_expert_id}:" + f"{tensor_name}" ) - if dst_slot is not None: - run_start = previous = dst_slot + return zlib.crc32(message_identity.encode("utf-8")) & 0x7FFFFFFF diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index afba43a7ad..8f53d70956 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -42,9 +42,10 @@ ) from lightllm.common.eplb_utils import extract_eplb_expert_tensors from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + EPLBTransferInfo, + ExpertTensorBuffer, PinnedMemoryEPLBTransfer, - TransferStep, - _commit_staging_rows, + TransferStatus, build_transfer_plan, ) @@ -540,10 +541,12 @@ def test_transfer_plan_respects_explicit_target_slots(): current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) target = torch.tensor([[5, 4], [7, 6], [1, 0], [3, 2]]) - plan = build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) + plan = build_transfer_plan(current, target, 3, num_logical_experts=8, world_size=4, node_world_size=2) - assert len(plan) == target.numel() - assert {(step.dst_rank, step.dst_slot) for step in plan} == {(rank, slot) for rank in range(4) for slot in range(2)} + assert all(info.layer_index == 3 for info in plan) + assert {(info.dest_rank, info.source_logical_expert_id) for info in plan} == { + (rank, int(target[rank, slot])) for rank in range(4) for slot in range(2) + } def test_manager_collects_logical_route_counters(): @@ -569,6 +572,9 @@ def test_manager_collects_logical_route_counters(): for counter in counters ] manager._eplb_impls = [weight.fuse_moe_impl for weight in manager.weights] + manager._weights = manager.weights + for layer_num, weight in enumerate(manager._weights): + weight.layer_num_ = layer_num manager.num_logical_experts = 2 samples = manager._collect_local_load() @@ -774,6 +780,9 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c for _ in range(3) ] manager._eplb_impls = [weight.fuse_moe_impl for weight in manager.weights] + manager._weights = manager.weights + for layer_num, weight in enumerate(manager._weights): + weight.layer_num_ = layer_num manager.global_rank = 1 manager.world_size = 4 manager.node_world_size = 2 @@ -787,7 +796,7 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c [ [[3], [0], [1], [2]], [[2], [3], [0], [1]], - [[1], [2], [3], [0]], + [[2], [3], [0], [1]], ], dtype=torch.int64, ) @@ -819,7 +828,7 @@ def build_maps_for_layers(*args, **kwargs): assert calls == [(2, 4, 2)] metadata = result["metadata"] assert 1 not in metadata - assert [layer_index for layer_index, _plan in result["layer_plans"]] == [0, 2] + assert {info.layer_index for info in result["transfer_infos"]} == {0, 2} for layer_index in (0, 2): item = metadata[layer_index] expected = build_logical_to_physical_map( @@ -1181,11 +1190,11 @@ def test_transfer_plan_uses_existing_rows_and_prefers_local_node_replicas(): target = current.clone() target[0, 0] = 6 # primary r3, but r1 replica is on r0's node. target[2, 1] = 4 # primary r2 is local to destination r2. - plan = build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) - by_dst = {(step.dst_rank, step.dst_slot): step for step in plan} - assert len(by_dst) == 2 - assert by_dst[0, 0] == TransferStep(0, 0, 1, 2) - assert by_dst[2, 1] == TransferStep(2, 1, 2, 0) + plan = build_transfer_plan(current, target, 5, num_logical_experts=8, world_size=4, node_world_size=2) + assert plan == [ + EPLBTransferInfo(1, 5, 6, 0), + EPLBTransferInfo(2, 5, 4, 2), + ] def test_transfer_plan_cross_node_and_stable_source_load_tie_break(): @@ -1193,131 +1202,99 @@ def test_transfer_plan_cross_node_and_stable_source_load_tie_break(): target = current.clone() target[0, 0] = 4 target[0, 1] = 4 - first = build_transfer_plan(current, target, 8, 4, 2) - second = build_transfer_plan(current, target, 8, 4, 2) + first = build_transfer_plan(current, target, 5, 8, 4, 2) + second = build_transfer_plan(current, target, 5, 8, 4, 2) assert first == second - selected = [step for step in first if step.dst_rank == 0] - assert [(step.src_rank, step.src_local_row) for step in selected] == [ - (2, 0), - (3, 2), + assert first == [ + EPLBTransferInfo(2, 5, 4, 0), + EPLBTransferInfo(3, 5, 4, 0), ] -def test_extract_expert_tensors_includes_weight_and_scale_in_order(): +def test_p2p_message_tag_is_stable_and_identifies_transfer_tensor(): + transfer = object.__new__(PinnedMemoryEPLBTransfer) + transfer.transfer_info = EPLBTransferInfo(1, 5, 4, 0) + weight_tag = transfer._build_p2p_message_tag("w13.weight") + + assert weight_tag == transfer._build_p2p_message_tag("w13.weight") + assert 0 <= weight_tag <= 0x7FFFFFFF + assert weight_tag != transfer._build_p2p_message_tag("w13.weight_scale") + transfer.transfer_info = EPLBTransferInfo(1, 5, 6, 0) + assert weight_tag != transfer._build_p2p_message_tag("w13.weight") + + +def test_extract_expert_tensors_includes_quantization_metadata_in_order(): class Pack: - def __init__(self, offset, scale=True): + def __init__(self, offset, scale=True, zero_point=True): self.weight = torch.full((3, 2), offset) self.weight_scale = torch.full((3, 1), offset + 1) if scale else None + self.weight_zero_point = torch.full((3, 1), offset + 2) if zero_point else None - weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False)})() + weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False, zero_point=False)})() tensors = extract_eplb_expert_tensors(weight) assert [name for name, _ in tensors] == [ "w13.weight", "w13.weight_scale", + "w13.weight_zero_point", "w2.weight", ] -def test_commit_staging_rows_only_overwrites_redundant_rows(): +def test_manager_commits_completed_transfer_rows_by_target_logical_expert(): live = torch.arange(20).reshape(5, 4) - staging = torch.full((2, 4), -1) - _commit_staging_rows( - live, - staging, - num_experts_per_rank=3, - changed_dst_slots=(0, 1), - ) - assert torch.equal(live[:3], torch.arange(12).reshape(3, 4)) - assert torch.equal(live[3:], staging) - - -def test_commit_staging_rows_preserves_unchanged_destination_slots(): - live = torch.arange(28).reshape(7, 4) - staging = torch.tensor([[-1, -1, -1, -1], [-2, -2, -2, -2], [-3, -3, -3, -3], [-4, -4, -4, -4]]) - original = live.clone() - - _commit_staging_rows( - live, - staging, - num_experts_per_rank=3, - changed_dst_slots=(3, 1), - ) - - assert torch.equal(live[:4], original[:4]) - assert torch.equal(live[4], staging[1]) - assert torch.equal(live[5], original[5]) - assert torch.equal(live[6], staging[3]) - - -def test_commit_staging_rows_merges_contiguous_changed_slots(): - copies = [] - - class View: - def __init__(self, owner, start, length): - self.owner = owner - self.start = start - self.length = length - - def copy_(self, source, **_kwargs): - copies.append( - ( - self.owner, - self.start, - self.length, - source.owner, - source.start, - source.length, - ) - ) - - class Tensor: - def __init__(self, owner, rows): - self.owner = owner - self.shape = (rows,) - - def narrow(self, _dim, start, length): - return View(self.owner, start, length) + original_primary = live[:3].clone() + local_expert_ids = [0, 1, 2, 4, 5] + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.global_rank = 0 + manager.num_primary_experts_per_rank = 3 + manager.target_placement = torch.tensor([[[7, 6]]]) + manager._eplb_impls = [SimpleNamespace(local_logics_expert_ids_list=local_expert_ids)] + manager._commit_layer_metadata = lambda _layer: None + manager.completed_transfers = [ + SimpleNamespace( + transfer_info=EPLBTransferInfo(1, 0, 7, 0), + tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -7))], + ), + SimpleNamespace( + transfer_info=EPLBTransferInfo(2, 0, 6, 0), + tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -6))], + ), + ] - _commit_staging_rows( - Tensor("live", 20), - Tensor("staging", 4), - num_experts_per_rank=10, - changed_dst_slots=(3, 1, 2), - ) + manager._commit_transferred_layer(0) - assert copies == [("live", 11, 3, "staging", 1, 3)] + assert torch.equal(live[:3], original_primary) + assert torch.equal(live[3], torch.full((4,), -7)) + assert torch.equal(live[4], torch.full((4,), -6)) + assert local_expert_ids == [0, 1, 2, 7, 6] -def test_manager_inflight_ready_gate_commits_one_layer_and_propagates_worker_error( - monkeypatch, -): +def test_manager_inflight_ready_gate_commits_one_layer(monkeypatch): class Transfer: - def __init__(self): - self.ready = 0 - self.commits = [] - self.finished = 0 - - def ready_layer(self): - return self.ready - - def commit(self, layer, post_copy=None): - assert self.ready == layer - self.ready = None - self.commits.append(layer) - if post_copy is not None: - post_copy() + def __init__(self, transfer_info): + self.transfer_info = transfer_info + self.status = TransferStatus.SUCCEEDED - def finish(self): - self.finished += 1 + def is_finished(self): + return True + info0 = EPLBTransferInfo(0, 0, 2, 1) + info0b = EPLBTransferInfo(1, 0, 3, 0) + info1 = EPLBTransferInfo(0, 1, 4, 1) manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.transfer = Transfer() + manager.active_transfer = Transfer(info0) manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager.in_flight_layers = [0, 1] + manager._control_status_buffer = torch.empty(1, dtype=torch.int32) + manager.in_flight_transfers = [info0, info0b, info1] + manager.completed_transfers = [] committed, finished = [], [] - manager._commit_layer_metadata = committed.append + manager._commit_transferred_layer = committed.append manager._finish_rebalance = lambda: finished.append(True) + + def start_next_transfer(): + manager.active_transfer = Transfer(manager.in_flight_transfers[0]) + + manager._start_next_transfer = start_next_transfer operations = [] class CurrentStream: @@ -1333,81 +1310,33 @@ def set_global_ready(count): monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(0)) manager._poll_in_flight() - assert manager.transfer.commits == [] + assert committed == [] monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) manager._poll_in_flight() - assert manager.transfer.commits == [0] + assert committed == [] + assert operations == [] + assert not finished + + manager._poll_in_flight() assert committed == [0] assert operations == [("wait", overlap_stream)] - assert not finished - manager.transfer.ready = 1 manager._poll_in_flight() - assert manager.transfer.commits == [0, 1] assert committed == [0, 1] assert finished == [True] - assert manager.transfer.finished == 1 - - manager.in_flight_layers = [2] - manager.transfer.ready = 9 - with pytest.raises(RuntimeError, match="does not match expected"): - manager._poll_in_flight() - - class BrokenTransfer: - def ready_layer(self): - raise RuntimeError("boom") - - manager.transfer = BrokenTransfer() - encoded_statuses = [] - - def retain_local_error(tensor, **_kwargs): - encoded_statuses.append(int(tensor.item())) - - monkeypatch.setattr(manager_module.dist, "all_reduce", retain_local_error) - with pytest.raises(RuntimeError, match="EPLB transfer worker failed on this rank") as exc_info: - manager._poll_in_flight() - assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "boom" - assert encoded_statuses == [manager_module.EPLB_CONTROL_ERROR] - - -def test_manager_inflight_remote_worker_error_does_not_commit(monkeypatch): - class Transfer: - def __init__(self): - self.commits = [] - - def ready_layer(self): - return 0 - - def commit(self, *args): - self.commits.append(args) - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.transfer = Transfer() - manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager.in_flight_layers = [0] - manager._commit_layer_metadata = lambda _layer: None - manager._finish_rebalance = lambda: None - statuses = [] - - def remote_error(tensor, **_kwargs): - statuses.append(int(tensor.item())) - tensor.fill_(manager_module.EPLB_CONTROL_ERROR) - - monkeypatch.setattr(manager_module.dist, "all_reduce", remote_error) - with pytest.raises(RuntimeError, match="EPLB transfer worker failed on another rank"): + expected = EPLBTransferInfo(0, 2, 5, 1) + manager.in_flight_transfers = [expected] + manager.active_transfer = Transfer(EPLBTransferInfo(1, 2, 5, 1)) + with pytest.raises(RuntimeError, match="does not match"): manager._poll_in_flight() - assert statuses == [1] - assert manager.transfer.commits == [] def test_evaluation_ready_gate_propagates_local_and_remote_errors(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager._control_status_buffer = torch.empty(1, dtype=torch.int32) manager._evaluation = Future() manager._evaluation.set_exception(RuntimeError("evaluation boom")) statuses = [] @@ -1441,26 +1370,16 @@ def test_manager_inflight_commit_orders_live_weights_between_overlap_forwards( monkeypatch, ): class Transfer: - def __init__(self, live, staging): - self.live = live - self.staging = staging - self.ready = 0 - - def ready_layer(self): - return self.ready - - def commit(self, layer, post_copy=None): - assert self.ready == layer - self.ready = None - self.live.copy_(self.staging, non_blocking=True) - if post_copy is not None: - post_copy() - - def finish(self): - pass + def __init__(self, live, received, transfer_info): + self.tensor_buffers = [ExpertTensorBuffer("weight", live, received)] + self.transfer_info = transfer_info + self.status = TransferStatus.SUCCEEDED + + def is_finished(self): + return True live = torch.tensor([1.0], device="cuda") - staging = torch.tensor([2.0], device="cuda") + received = torch.tensor(2.0, pin_memory=True) previous_read = torch.empty_like(live) next_read = torch.empty_like(live) source_stream = torch.cuda.Stream(device=live.device) @@ -1469,10 +1388,16 @@ def finish(self): original_overlap_stream = g_infer_context.overlap_stream manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.transfer = Transfer(live, staging) + transfer_info = EPLBTransferInfo(0, 0, 2, 0) + manager.active_transfer = Transfer(live, received, transfer_info) manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager.in_flight_layers = [0] + manager._control_status_buffer = torch.empty(1, dtype=torch.int32) + manager.in_flight_transfers = [transfer_info] + manager.completed_transfers = [] + manager.num_primary_experts_per_rank = 0 + manager.global_rank = 0 + manager.target_placement = torch.tensor([[[2]]]) + manager._eplb_impls = [SimpleNamespace(local_logics_expert_ids_list=[1])] manager._commit_layer_metadata = lambda _layer: None manager._finish_rebalance = lambda: None monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) @@ -1498,7 +1423,7 @@ def finish(self): def test_manager_inflight_step_does_not_poll(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_layers = [0] + manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1)] manager._evaluation = None calls = [] manager._poll_in_flight = lambda: calls.append("poll") @@ -1508,7 +1433,7 @@ def test_manager_inflight_step_does_not_poll(): def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_layers = [] + manager.in_flight_transfers = [] manager._evaluation = None manager.prefill_steps = 0 manager.step_interval = 3 @@ -1547,7 +1472,7 @@ def test_manager_restart_collection_resets_counters_and_arms_next_window(): def test_manager_poll_advances_inflight_before_evaluation(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_layers = [0] + manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1)] manager._evaluation = Future() calls = [] manager._poll_in_flight = lambda: calls.append("inflight") @@ -1560,10 +1485,10 @@ def test_manager_poll_advances_inflight_before_evaluation(): def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_layers = [] + manager.in_flight_transfers = [] manager._evaluation = Future() manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager._control_status_buffer = torch.empty(1, dtype=torch.int32) calls = [] manager._finish_evaluation = lambda: calls.append("evaluation") @@ -1604,20 +1529,7 @@ def test_mode_backend_eplb_hooks_are_noops_when_disabled(): assert calls == ["prefill"] -def test_pinned_transfer_groups_sources_in_collective_order(): - plan = [ - TransferStep(0, 2, 1, 3), - TransferStep(0, 0, 0, 1), - TransferStep(1, 1, 1, 3), - ] - - grouped = PinnedMemoryEPLBTransfer._group_steps_by_source(plan) - - assert [source for source, _steps in grouped] == [(0, 1), (1, 3)] - assert grouped[1][1] == [plan[0], plan[2]] - - -def test_pinned_transfer_copies_source_row_through_cpu_buffer(monkeypatch): +def test_pinned_transfer_copies_source_row_and_sends_to_destination(monkeypatch): class Stream: def __init__(self): self.synchronize_count = 0 @@ -1626,84 +1538,151 @@ def synchronize(self): self.synchronize_count += 1 transfer = object.__new__(PinnedMemoryEPLBTransfer) - transfer.global_rank = 0 - transfer.transfer_group = object() - transfer._copy_stream = Stream() - transfer.live = [[("weight", torch.tensor([[1.0, 2.0], [3.0, 4.0]]))]] - transfer.pinned_rows = [("weight", torch.empty(2))] - transfer.staging = [("weight", torch.zeros((2, 2)))] - broadcasts = [] + transfer._device = "cuda:0" + transfer._is_source_rank = True + transfer._is_destination_rank = False + transfer._p2p_group = object() + transfer.transfer_info = EPLBTransferInfo(0, 0, 5, 1) + transfer._device_to_host_stream = Stream() + transfer._local_logical_expert_ids = [4, 5] + transfer.tensor_buffers = [ + ExpertTensorBuffer( + "weight", + torch.tensor([[1.0, 2.0], [3.0, 4.0]]), + torch.empty(2), + ) + ] + transfer.status = TransferStatus.RUNNING + sends = [] + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: nullcontext()) monkeypatch.setattr( transfer_module.dist, - "broadcast", - lambda tensor, src, group: broadcasts.append((tensor.clone(), src, group)), - ) - - transfer._copy_layer( - 0, - [ - TransferStep(dst_rank=0, dst_slot=1, src_rank=0, src_local_row=1), - TransferStep(dst_rank=1, dst_slot=0, src_rank=0, src_local_row=1), - ], + "send", + lambda tensor, dst, group, tag: sends.append((tensor.clone(), dst, group, tag)), ) - assert len(broadcasts) == 1 - assert torch.equal(broadcasts[0][0], torch.tensor([3.0, 4.0])) - assert broadcasts[0][1:] == (0, transfer.transfer_group) - assert torch.equal(transfer.staging[0][1][1], torch.tensor([3.0, 4.0])) - assert torch.count_nonzero(transfer.staging[0][1][0]) == 0 - assert transfer._copy_stream.synchronize_count == 2 + transfer._run_transfer() + assert transfer.status is TransferStatus.SUCCEEDED + assert len(sends) == 1 + assert torch.equal(sends[0][0], torch.tensor([3.0, 4.0])) + expected_tag = transfer._build_p2p_message_tag("weight") + assert sends[0][1:] == (1, transfer._p2p_group, expected_tag) + assert torch.equal(transfer.tensor_buffers[0].pinned_row, torch.tensor([3.0, 4.0])) + assert transfer._device_to_host_stream.synchronize_count == 1 -def test_pinned_transfer_waits_for_each_layer_commit_before_reusing_staging( - monkeypatch, -): - class Event: - def __init__(self): - self.synchronize_count = 0 +def test_pinned_transfer_skips_p2p_for_local_destination(monkeypatch): + class Stream: def synchronize(self): - self.synchronize_count += 1 - - def record(self, _stream): pass transfer = object.__new__(PinnedMemoryEPLBTransfer) - transfer.device = "cuda:0" - transfer.global_rank = 0 - transfer.live = [[], []] - transfer.staging = [] - transfer.num_experts_per_rank = 0 - transfer._release = threading.Event() - transfer._release.set() - transfer._consumed_event = Event() - transfer._consumed_recorded = False - transfer._ready = None - transfer._ready_lock = threading.Lock() - transfer._error = None - transfer._thread = None - copied = [] - transfer._copy_layer = lambda layer_index, _plan: copied.append(layer_index) + transfer._device = "cuda:0" + transfer._is_source_rank = True + transfer._is_destination_rank = True + transfer._p2p_group = object() + transfer.transfer_info = EPLBTransferInfo(0, 0, 5, 0) + transfer._device_to_host_stream = Stream() + transfer._local_logical_expert_ids = [5] + transfer.tensor_buffers = [ + ExpertTensorBuffer( + "weight", + torch.tensor([[3.0, 4.0]]), + torch.empty(2), + ) + ] + transfer.status = TransferStatus.RUNNING + p2p_calls = [] monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) - monkeypatch.setattr(transfer_module.torch.cuda, "current_stream", lambda: object()) + monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: nullcontext()) + monkeypatch.setattr(transfer_module.dist, "send", lambda *_args, **_kwargs: p2p_calls.append("send")) + monkeypatch.setattr(transfer_module.dist, "recv", lambda *_args, **_kwargs: p2p_calls.append("recv")) - transfer.start([(0, []), (1, [])]) - deadline = time.monotonic() + 2 - while transfer.ready_layer() is None and time.monotonic() < deadline: - time.sleep(0.001) - assert transfer.ready_layer() == 0 - assert copied == [0] + transfer._run_transfer() + + assert transfer.is_finished() + assert p2p_calls == [] + assert torch.equal(transfer.tensor_buffers[0].pinned_row, torch.tensor([3.0, 4.0])) + + +def test_pinned_transfer_exits_process_on_failure(monkeypatch): + transfer = object.__new__(PinnedMemoryEPLBTransfer) + transfer._device = "cuda:0" + transfer._is_source_rank = False + transfer._is_destination_rank = True + transfer.transfer_info = EPLBTransferInfo(1, 0, 3, 0) + transfer._p2p_group = object() + transfer.tensor_buffers = [ + ExpertTensorBuffer( + "weight", + torch.empty((1, 1)), + torch.empty(1), + ) + ] + transfer.status = TransferStatus.RUNNING + logged_messages = [] + exit_codes = [] + + def fail_recv(*_args, **_kwargs): + raise RuntimeError("recv failed") + + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + monkeypatch.setattr(transfer_module.dist, "recv", fail_recv) + monkeypatch.setattr(transfer_module.logger, "exception", logged_messages.append) + monkeypatch.setattr(transfer_module.os, "_exit", exit_codes.append) + + transfer._run_transfer() + + assert logged_messages == ["EPLB transfer failed"] + assert exit_codes == [1] + assert not transfer.is_finished() + + +def test_pinned_transfer_is_single_use_and_exposes_pinned_rows(monkeypatch): + class Stream: + def synchronize(self): + pass + + transfer = object.__new__(PinnedMemoryEPLBTransfer) + transfer._device = "cuda:0" + transfer._is_source_rank = False + transfer._is_destination_rank = True + transfer.transfer_info = EPLBTransferInfo(1, 0, 3, 0) + transfer._p2p_group = object() + transfer._device_to_host_stream = Stream() + transfer.tensor_buffers = [ + ExpertTensorBuffer( + "weight", + torch.empty((1, 1)), + torch.tensor([3.0]), + ) + ] + transfer.status = TransferStatus.IDLE + transfer._transfer_thread = threading.Thread(target=transfer._run_transfer, daemon=True) + receives = [] + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + monkeypatch.setattr( + transfer_module.dist, + "recv", + lambda tensor, src, group, tag: receives.append((tensor, src, group, tag)), + ) - transfer.commit(0) + assert not transfer.is_finished() + transfer.start() deadline = time.monotonic() + 2 - while transfer.ready_layer() is None and time.monotonic() < deadline: + while transfer.status is TransferStatus.RUNNING and time.monotonic() < deadline: time.sleep(0.001) - assert transfer.ready_layer() == 1 - assert copied == [0, 1] - assert transfer._consumed_event.synchronize_count == 1 - transfer.commit(1) - transfer.finish() + assert transfer.status is TransferStatus.SUCCEEDED + assert transfer.is_finished() + assert torch.equal(transfer.tensor_buffers[0].pinned_row, torch.tensor([3.0])) + assert len(receives) == 1 + assert receives[0][0] is transfer.tensor_buffers[0].pinned_row + expected_tag = transfer._build_p2p_message_tag("weight") + assert receives[0][1:] == (1, transfer._p2p_group, expected_tag) + with pytest.raises(AssertionError, match="already been started"): + transfer.start() def test_manager_constructs_pinned_memory_transfer(monkeypatch): @@ -1712,6 +1691,7 @@ def test_manager_constructs_pinned_memory_transfer(monkeypatch): (), { "n_routed_experts": 4, + "layer_num_": 0, "fuse_moe_impl": _test_moe_impl( eplb=True, num_logical_experts=4, @@ -1721,7 +1701,8 @@ def test_manager_constructs_pinned_memory_transfer(monkeypatch): ), }, )() - transfer = object() + transfer_starts = [] + transfer = SimpleNamespace(start=lambda: transfer_starts.append(True)) groups = [object(), object(), object()] new_group_calls = [] monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) @@ -1740,19 +1721,23 @@ def new_group(*args, **kwargs): monkeypatch.setattr( manager_module, "PinnedMemoryEPLBTransfer", - lambda weights, group, rank: (transfer_calls.append((weights, group, rank)) or transfer), + lambda weights, group, rank, info: (transfer_calls.append((weights, group, rank, info)) or transfer), ) logs = [] monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) manager = manager_module.EPLBManager(type("Model", (), {})()) - assert manager.transfer is transfer + assert manager.active_transfer is None assert ( manager.evaluation_group, manager.control_group, manager.transfer_group, ) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 - assert transfer_calls == [([weight], groups[2], 0)] + transfer_info = EPLBTransferInfo(0, 0, 2, 1) + manager.in_flight_transfers = [transfer_info] + manager._start_next_transfer() + assert transfer_calls == [([weight], groups[2], 0, transfer_info)] + assert transfer_starts == [True] assert manager.planner.rebalance_gain_threshold == 0.07 assert manager.next_evaluation_step == manager.step_interval assert "planner=GreedyEPLBPlanner" in logs[0] diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 372f1c10ce..758da5d183 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -12,6 +12,7 @@ from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( PinnedMemoryEPLBTransfer, + TransferStatus, build_transfer_plan, ) @@ -25,11 +26,13 @@ def __init__(self, weight, weight_scale): class _FakeWeight: def __init__(self, rank, layer_index): + self.layer_num_ = layer_index + logical_ids = ([0, 1, 2], [2, 3, 0])[rank] self.fuse_moe_impl = SimpleNamespace( num_primary_experts_per_rank=2, num_redundant_experts_per_rank=1, + local_logics_expert_ids_list=list(logical_ids), ) - logical_ids = ([0, 1, 2], [2, 3, 0])[rank] self.w13 = self._pack(logical_ids, layer_index, 0) self.w2 = self._pack(logical_ids, layer_index, 10) @@ -52,16 +55,15 @@ def _free_port(): return port -def _wait_for_ready_layer(transfer, control_group): +def _wait_for_transfer(transfer, control_group): deadline = time.monotonic() + 30 while time.monotonic() < deadline: - ready_layer = transfer.ready_layer() - ready_count = torch.tensor([int(ready_layer is not None)], dtype=torch.int32) + ready_count = torch.tensor([int(transfer.status is TransferStatus.SUCCEEDED)], dtype=torch.int32) dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=control_group) if int(ready_count.item()) == 1: - return ready_layer + return time.sleep(0.001) - raise TimeoutError("EPLB transfer worker did not publish a globally ready layer") + raise TimeoutError("EPLB transfer worker did not finish globally") def _worker(rank, port): @@ -73,33 +75,37 @@ def _worker(rank, port): transfer_group = dist.new_group([0, 1], backend="gloo") weights = [_FakeWeight(rank, layer_index) for layer_index in range(2)] - transfer = PinnedMemoryEPLBTransfer(weights, transfer_group, rank) - assert all(row.is_pinned() for _, row in transfer.pinned_rows) current = torch.tensor([[2], [0]]) target = torch.tensor([[3], [1]]) - plan = build_transfer_plan(current, target, num_logical_experts=4, world_size=2, node_world_size=2) - layer_plans = [(0, plan), (1, plan)] - transfer.start(layer_plans) - for expected_layer in range(2): - layer_index = _wait_for_ready_layer(transfer, control_group) - assert layer_index == expected_layer - expected_expert = 3 if rank == 0 else 1 - expected_w13 = expected_layer * 100 + expected_expert - expected_w2 = expected_w13 + 10 - assert torch.all(transfer.staging[0][1][0] == expected_w13) - assert torch.all(transfer.staging[1][1][0] == expected_w13 + 0.5) - assert torch.all(transfer.staging[2][1][0] == expected_w2) - assert torch.all(transfer.staging[3][1][0] == expected_w2 + 0.5) - transfer.commit(layer_index) - - transfer.finish() - torch.cuda.synchronize() - for layer_index, weight in enumerate(weights): - expected_expert = 3 if rank == 0 else 1 - expected_w13 = layer_index * 100 + expected_expert - assert torch.all(weight.w13.weight[2] == expected_w13) - assert torch.all(weight.w2.weight[2] == expected_w13 + 10) + transfer_infos = build_transfer_plan( + current, + target, + expected_layer, + num_logical_experts=4, + world_size=2, + node_world_size=2, + ) + for transfer_info in transfer_infos: + transfer = PinnedMemoryEPLBTransfer(weights, transfer_group, rank, transfer_info) + assert all(buffer.pinned_row.is_pinned() for buffer in transfer.tensor_buffers) + transfer.start() + _wait_for_transfer(transfer, control_group) + if transfer_info.dest_rank == rank: + expected_expert = 3 if rank == 0 else 1 + expected_w13 = expected_layer * 100 + expected_expert + expected_w2 = expected_w13 + 10 + assert [buffer.name for buffer in transfer.tensor_buffers] == [ + "w13.weight", + "w13.weight_scale", + "w2.weight", + "w2.weight_scale", + ] + pinned_rows = [buffer.pinned_row for buffer in transfer.tensor_buffers] + assert torch.all(pinned_rows[0][0] == expected_w13) + assert torch.all(pinned_rows[1][0] == expected_w13 + 0.5) + assert torch.all(pinned_rows[2][0] == expected_w2) + assert torch.all(pinned_rows[3][0] == expected_w2 + 0.5) dist.destroy_process_group() From a0a970b12ba8fe1cf78c7a15acb61ab28849de1e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 07:36:28 +0000 Subject: [PATCH 28/72] refactor(eplb): simplify expert transfer planning --- .../model_infer/mode_backend/eplb_manager.py | 12 +- .../model_infer/mode_backend/eplb_transfer.py | 131 ++++++------------ unit_tests/common/fused_moe/test_eplb.py | 79 +++++------ .../fused_moe/test_eplb_transfer_gpu.py | 5 +- 4 files changed, 88 insertions(+), 139 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 09701f7b83..a95a3f9c33 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -27,7 +27,6 @@ from lightllm.utils.dist_utils import ( get_global_rank, get_global_world_size, - get_node_world_size, ) from lightllm.utils.envs_utils import ( get_eplb_rebalance_gain_threshold, @@ -51,7 +50,6 @@ def __init__(self, model: TpPartBaseModel) -> None: self._weights: List[FusedMoeWeight] = weights self.global_rank: int = get_global_rank() self.world_size: int = get_global_world_size() - self.node_world_size: int = get_node_world_size() self._eplb_impls = [weight.fuse_moe_impl for weight in weights] routed = {impl.n_routed_experts for impl in self._eplb_impls} redundant = {impl.num_redundant_experts_per_rank for impl in self._eplb_impls} @@ -183,16 +181,16 @@ def _build_rebalance_data(self, result: Dict[str, Any]) -> Tuple[Dict[int, torch dtype=torch.int32, ) for changed_layer_offset, layer_index in enumerate(changed_layer_indices): - target_layer_placement = torch.tensor(result["placement"][layer_index], dtype=torch.int64) + current_layer_placement = self.current_placement[layer_index].tolist() + target_layer_placement = result["placement"][layer_index] metadata_by_layer[layer_index] = logical_to_physical_maps[changed_layer_offset] planned_transfers.extend( build_transfer_plan( - self.current_placement[layer_index], + current_layer_placement, target_layer_placement, layer_index, self.num_logical_experts, self.world_size, - self.node_world_size, ) ) return metadata_by_layer, planned_transfers @@ -324,10 +322,8 @@ def _commit_transferred_layer(self, layer_index: int) -> None: transfer_info: EPLBTransferInfo = transfer.transfer_info if transfer_info.dest_rank != self.global_rank: continue - destination_slot_index: int = target_redundant_expert_ids.index(transfer_info.source_logical_expert_id) - destination_local_expert_index: int = self.num_primary_experts_per_rank + destination_slot_index for tensor_buffer in transfer.tensor_buffers: - tensor_buffer.live_tensor[destination_local_expert_index].copy_(tensor_buffer.pinned_row) + tensor_buffer.live_tensor[transfer_info.dest_local_expert_index].copy_(tensor_buffer.pinned_row) local_expert_ids: List[int] = self._eplb_impls[layer_index].local_logics_expert_ids_list local_expert_ids[self.num_primary_experts_per_rank :] = target_redundant_expert_ids diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index 05f0cbf0ff..0793a3f989 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -31,22 +31,15 @@ class EPLBTransferInfo: """单个逻辑专家的一次传输描述。 ``layer_index`` 是专家权重在 EPLB 层列表中的下标。源 rank 使用 - ``source_logical_expert_id`` 定位当前本地物理行,目标 rank 从 - ``tensor_buffers`` 中读取传输完成的 pinned memory 数据。 + ``source_logical_expert_id`` 定位当前本地物理行;目标 rank 将收到的 + pinned memory 数据写入 ``dest_local_expert_index`` 指定的本地物理行。 """ source_rank: int layer_index: int source_logical_expert_id: int dest_rank: int - - -@dataclass(frozen=True) -class _ExpertSource: - """一个逻辑专家当前可用的物理副本位置。""" - - rank: int - local_expert_index: int + dest_local_expert_index: int class TransferStatus(Enum): @@ -57,82 +50,6 @@ class TransferStatus(Enum): SUCCEEDED = "succeeded" -def build_transfer_plan( - current_placement: torch.Tensor, - target_placement: torch.Tensor, - layer_index: int, - num_logical_experts: int, - world_size: int, - node_world_size: int, -) -> List[EPLBTransferInfo]: - """根据新旧冗余专家分布生成确定性的传输计划。 - - ``current_placement`` 和 ``target_placement`` 只描述冗余槽位,形状均为 - ``[world_size, num_redundant_slots]``。固定主专家不在这两个张量中,但始终 - 可以作为数据源。返回结果只包含发生变化的目标槽位所需专家,每个 - :class:`EPLBTransferInfo` 只描述一个逻辑专家的传输。 - - 为同一个逻辑专家选择数据源时,依次考虑: - - 1. 优先使用目标 rank 所在节点上的已有副本,避免跨节点传输; - 2. 均衡各个源 rank 承担的传输次数; - 3. 使用 rank 和本地行号做稳定排序,保证所有 rank 生成一致结果。 - """ - assert ( - tuple(current_placement.shape) - == tuple(target_placement.shape) - == ( - world_size, - current_placement.shape[1], - ) - ) - num_primary_experts_per_rank = num_logical_experts // world_size - current_placement_by_rank: List[List[int]] = current_placement.tolist() - target_placement_by_rank: List[List[int]] = target_placement.tolist() - - # 每个逻辑专家的主专家行永远存在,因此先把它加入候选源;当前仍存在的 - # 冗余副本也可以作为源,这样目标 rank 有机会直接使用同节点副本。 - source_candidates_by_expert: List[List[_ExpertSource]] = [] - for logical_expert_id in range(num_logical_experts): - primary_rank, primary_local_expert_index = divmod(logical_expert_id, num_primary_experts_per_rank) - source_candidates_by_expert.append([_ExpertSource(primary_rank, primary_local_expert_index)]) - for rank, redundant_expert_ids in enumerate(current_placement_by_rank): - for redundant_slot_index, logical_expert_id in enumerate(redundant_expert_ids): - source_candidates_by_expert[logical_expert_id].append( - _ExpertSource( - rank=rank, - local_expert_index=num_primary_experts_per_rank + redundant_slot_index, - ) - ) - - num_transfers_by_source_rank: List[int] = [0] * world_size - transfer_infos: List[EPLBTransferInfo] = [] - for destination_rank in range(world_size): - for destination_slot_index, logical_expert_id in enumerate(target_placement_by_rank[destination_rank]): - if logical_expert_id == current_placement_by_rank[destination_rank][destination_slot_index]: - continue - source = min( - source_candidates_by_expert[logical_expert_id], - key=lambda candidate: ( - candidate.rank // node_world_size != destination_rank // node_world_size, - num_transfers_by_source_rank[candidate.rank], - candidate.rank, - candidate.local_expert_index, - ), - ) - num_transfers_by_source_rank[source.rank] += 1 - transfer_infos.append( - EPLBTransferInfo( - source_rank=source.rank, - layer_index=layer_index, - source_logical_expert_id=logical_expert_id, - dest_rank=destination_rank, - ) - ) - - return transfer_infos - - class PinnedMemoryEPLBTransfer: """在后台线程中传输一个逻辑专家的全部权重张量。 @@ -265,6 +182,48 @@ def _build_p2p_message_tag(self, tensor_name: str) -> int: f"{transfer_info.source_rank}:" f"{transfer_info.dest_rank}:" f"{transfer_info.source_logical_expert_id}:" + f"{transfer_info.dest_local_expert_index}:" f"{tensor_name}" ) return zlib.crc32(message_identity.encode("utf-8")) & 0x7FFFFFFF + + +def build_transfer_plan( + current_placement: Sequence[Sequence[int]], + target_placement: Sequence[Sequence[int]], + layer_index: int, + num_logical_experts: int, + world_size: int, +) -> List[EPLBTransferInfo]: + """生成一层中所有发生变化的冗余专家传输任务。 + + ``current_placement`` 和 ``target_placement`` 的形状均为 + ``[world_size, num_redundant_slots]``。每个逻辑专家按照无冗余布局连续分配 + 给各 rank;该固定主副本始终作为传输源,不再从已有冗余副本中选择数据源。 + """ + assert world_size > 0 + assert num_logical_experts % world_size == 0 + assert len(current_placement) == len(target_placement) == world_size + num_redundant_slots = len(current_placement[0]) + assert all(len(row) == num_redundant_slots for row in current_placement) + assert all(len(row) == num_redundant_slots for row in target_placement) + + num_primary_experts_per_rank = num_logical_experts // world_size + transfer_infos: List[EPLBTransferInfo] = [] + for destination_rank, (current_row, target_row) in enumerate(zip(current_placement, target_placement)): + for destination_slot_index, (current_expert_id, target_expert_id) in enumerate(zip(current_row, target_row)): + if target_expert_id == current_expert_id: + continue + assert 0 <= target_expert_id < num_logical_experts + source_rank = target_expert_id // num_primary_experts_per_rank + transfer_infos.append( + EPLBTransferInfo( + source_rank=source_rank, + layer_index=layer_index, + source_logical_expert_id=target_expert_id, + dest_rank=destination_rank, + dest_local_expert_index=num_primary_experts_per_rank + destination_slot_index, + ) + ) + + return transfer_infos diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 8f53d70956..49967bc6fa 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -538,14 +538,14 @@ def test_logical_to_physical_maps_for_layers_match_single_layer_api(current_rank def test_transfer_plan_respects_explicit_target_slots(): - current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) - target = torch.tensor([[5, 4], [7, 6], [1, 0], [3, 2]]) + current = [[4, 5], [6, 7], [0, 1], [2, 3]] + target = [[5, 4], [7, 6], [1, 0], [3, 2]] - plan = build_transfer_plan(current, target, 3, num_logical_experts=8, world_size=4, node_world_size=2) + plan = build_transfer_plan(current, target, 3, num_logical_experts=8, world_size=4) assert all(info.layer_index == 3 for info in plan) assert {(info.dest_rank, info.source_logical_expert_id) for info in plan} == { - (rank, int(target[rank, slot])) for rank in range(4) for slot in range(2) + (rank, target[rank][slot]) for rank in range(4) for slot in range(2) } @@ -724,7 +724,6 @@ def test_manager_evaluation_collective_sums_rank_loads(monkeypatch): manager._eplb_impls = [manager.weights[0].fuse_moe_impl] manager.global_rank = 2 manager.world_size = 4 - manager.node_world_size = 2 manager.step_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 @@ -785,7 +784,6 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c weight.layer_num_ = layer_num manager.global_rank = 1 manager.world_size = 4 - manager.node_world_size = 2 manager.step_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 @@ -1185,41 +1183,39 @@ def fused(**kwargs): assert all(call["w13"] is w13 and call["w2"] is w2 for call in captured) -def test_transfer_plan_uses_existing_rows_and_prefers_local_node_replicas(): - current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) - target = current.clone() - target[0, 0] = 6 # primary r3, but r1 replica is on r0's node. - target[2, 1] = 4 # primary r2 is local to destination r2. - plan = build_transfer_plan(current, target, 5, num_logical_experts=8, world_size=4, node_world_size=2) +def test_transfer_plan_always_uses_primary_expert_rank(): + current = [[4, 5], [6, 7], [0, 1], [2, 3]] + target = [[6, 5], [6, 7], [0, 4], [2, 3]] + plan = build_transfer_plan(current, target, 5, num_logical_experts=8, world_size=4) assert plan == [ - EPLBTransferInfo(1, 5, 6, 0), - EPLBTransferInfo(2, 5, 4, 2), + EPLBTransferInfo(3, 5, 6, 0, 2), + EPLBTransferInfo(2, 5, 4, 2, 3), ] -def test_transfer_plan_cross_node_and_stable_source_load_tie_break(): - current = torch.tensor([[0, 1], [2, 3], [4, 5], [4, 7]]) - target = current.clone() - target[0, 0] = 4 - target[0, 1] = 4 - first = build_transfer_plan(current, target, 5, 8, 4, 2) - second = build_transfer_plan(current, target, 5, 8, 4, 2) +def test_transfer_plan_uses_same_primary_source_for_repeated_expert(): + current = [[0, 1], [2, 3], [4, 5], [4, 7]] + target = [[4, 4], [2, 3], [4, 5], [4, 7]] + first = build_transfer_plan(current, target, 5, 8, 4) + second = build_transfer_plan(current, target, 5, 8, 4) assert first == second assert first == [ - EPLBTransferInfo(2, 5, 4, 0), - EPLBTransferInfo(3, 5, 4, 0), + EPLBTransferInfo(2, 5, 4, 0, 2), + EPLBTransferInfo(2, 5, 4, 0, 3), ] def test_p2p_message_tag_is_stable_and_identifies_transfer_tensor(): transfer = object.__new__(PinnedMemoryEPLBTransfer) - transfer.transfer_info = EPLBTransferInfo(1, 5, 4, 0) + transfer.transfer_info = EPLBTransferInfo(1, 5, 4, 0, 2) weight_tag = transfer._build_p2p_message_tag("w13.weight") assert weight_tag == transfer._build_p2p_message_tag("w13.weight") assert 0 <= weight_tag <= 0x7FFFFFFF assert weight_tag != transfer._build_p2p_message_tag("w13.weight_scale") - transfer.transfer_info = EPLBTransferInfo(1, 5, 6, 0) + transfer.transfer_info = EPLBTransferInfo(1, 5, 6, 0, 2) + assert weight_tag != transfer._build_p2p_message_tag("w13.weight") + transfer.transfer_info = EPLBTransferInfo(1, 5, 4, 0, 3) assert weight_tag != transfer._build_p2p_message_tag("w13.weight") @@ -1240,7 +1236,7 @@ def __init__(self, offset, scale=True, zero_point=True): ] -def test_manager_commits_completed_transfer_rows_by_target_logical_expert(): +def test_manager_commits_completed_transfer_rows_by_planned_destination_index(): live = torch.arange(20).reshape(5, 4) original_primary = live[:3].clone() local_expert_ids = [0, 1, 2, 4, 5] @@ -1252,11 +1248,11 @@ def test_manager_commits_completed_transfer_rows_by_target_logical_expert(): manager._commit_layer_metadata = lambda _layer: None manager.completed_transfers = [ SimpleNamespace( - transfer_info=EPLBTransferInfo(1, 0, 7, 0), + transfer_info=EPLBTransferInfo(1, 0, 7, 0, 3), tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -7))], ), SimpleNamespace( - transfer_info=EPLBTransferInfo(2, 0, 6, 0), + transfer_info=EPLBTransferInfo(2, 0, 6, 0, 4), tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -6))], ), ] @@ -1278,9 +1274,9 @@ def __init__(self, transfer_info): def is_finished(self): return True - info0 = EPLBTransferInfo(0, 0, 2, 1) - info0b = EPLBTransferInfo(1, 0, 3, 0) - info1 = EPLBTransferInfo(0, 1, 4, 1) + info0 = EPLBTransferInfo(0, 0, 2, 1, 2) + info0b = EPLBTransferInfo(1, 0, 3, 0, 2) + info1 = EPLBTransferInfo(0, 1, 4, 1, 2) manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.active_transfer = Transfer(info0) manager.control_group = object() @@ -1326,9 +1322,9 @@ def set_global_ready(count): assert committed == [0, 1] assert finished == [True] - expected = EPLBTransferInfo(0, 2, 5, 1) + expected = EPLBTransferInfo(0, 2, 5, 1, 2) manager.in_flight_transfers = [expected] - manager.active_transfer = Transfer(EPLBTransferInfo(1, 2, 5, 1)) + manager.active_transfer = Transfer(EPLBTransferInfo(1, 2, 5, 1, 2)) with pytest.raises(RuntimeError, match="does not match"): manager._poll_in_flight() @@ -1388,7 +1384,7 @@ def is_finished(self): original_overlap_stream = g_infer_context.overlap_stream manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - transfer_info = EPLBTransferInfo(0, 0, 2, 0) + transfer_info = EPLBTransferInfo(0, 0, 2, 0, 0) manager.active_transfer = Transfer(live, received, transfer_info) manager.control_group = object() manager._control_status_buffer = torch.empty(1, dtype=torch.int32) @@ -1423,7 +1419,7 @@ def is_finished(self): def test_manager_inflight_step_does_not_poll(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1)] + manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1, 2)] manager._evaluation = None calls = [] manager._poll_in_flight = lambda: calls.append("poll") @@ -1472,7 +1468,7 @@ def test_manager_restart_collection_resets_counters_and_arms_next_window(): def test_manager_poll_advances_inflight_before_evaluation(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1)] + manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1, 2)] manager._evaluation = Future() calls = [] manager._poll_in_flight = lambda: calls.append("inflight") @@ -1542,7 +1538,7 @@ def synchronize(self): transfer._is_source_rank = True transfer._is_destination_rank = False transfer._p2p_group = object() - transfer.transfer_info = EPLBTransferInfo(0, 0, 5, 1) + transfer.transfer_info = EPLBTransferInfo(0, 0, 5, 1, 2) transfer._device_to_host_stream = Stream() transfer._local_logical_expert_ids = [4, 5] transfer.tensor_buffers = [ @@ -1583,7 +1579,7 @@ def synchronize(self): transfer._is_source_rank = True transfer._is_destination_rank = True transfer._p2p_group = object() - transfer.transfer_info = EPLBTransferInfo(0, 0, 5, 0) + transfer.transfer_info = EPLBTransferInfo(0, 0, 5, 0, 1) transfer._device_to_host_stream = Stream() transfer._local_logical_expert_ids = [5] transfer.tensor_buffers = [ @@ -1612,7 +1608,7 @@ def test_pinned_transfer_exits_process_on_failure(monkeypatch): transfer._device = "cuda:0" transfer._is_source_rank = False transfer._is_destination_rank = True - transfer.transfer_info = EPLBTransferInfo(1, 0, 3, 0) + transfer.transfer_info = EPLBTransferInfo(1, 0, 3, 0, 2) transfer._p2p_group = object() transfer.tensor_buffers = [ ExpertTensorBuffer( @@ -1649,7 +1645,7 @@ def synchronize(self): transfer._device = "cuda:0" transfer._is_source_rank = False transfer._is_destination_rank = True - transfer.transfer_info = EPLBTransferInfo(1, 0, 3, 0) + transfer.transfer_info = EPLBTransferInfo(1, 0, 3, 0, 2) transfer._p2p_group = object() transfer._device_to_host_stream = Stream() transfer.tensor_buffers = [ @@ -1708,7 +1704,6 @@ def test_manager_constructs_pinned_memory_transfer(monkeypatch): monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) - monkeypatch.setattr(manager_module, "get_node_world_size", lambda: 2) monkeypatch.setattr(manager_module, "get_prefill_eplb_step_interval", lambda: 20) monkeypatch.setattr(manager_module, "get_eplb_rebalance_gain_threshold", lambda: 0.07) @@ -1733,7 +1728,7 @@ def new_group(*args, **kwargs): manager.transfer_group, ) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 - transfer_info = EPLBTransferInfo(0, 0, 2, 1) + transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) manager.in_flight_transfers = [transfer_info] manager._start_next_transfer() assert transfer_calls == [([weight], groups[2], 0, transfer_info)] diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 758da5d183..da05afb515 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -75,8 +75,8 @@ def _worker(rank, port): transfer_group = dist.new_group([0, 1], backend="gloo") weights = [_FakeWeight(rank, layer_index) for layer_index in range(2)] - current = torch.tensor([[2], [0]]) - target = torch.tensor([[3], [1]]) + current = [[2], [0]] + target = [[3], [1]] for expected_layer in range(2): transfer_infos = build_transfer_plan( current, @@ -84,7 +84,6 @@ def _worker(rank, port): expected_layer, num_logical_experts=4, world_size=2, - node_world_size=2, ) for transfer_info in transfer_infos: transfer = PinnedMemoryEPLBTransfer(weights, transfer_group, rank, transfer_info) From b133da6ab97ac05f37277dabc8c8ff7d8ec869aa Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 07:52:00 +0000 Subject: [PATCH 29/72] refactor(eplb): unify manager lifecycle step --- .../model_infer/mode_backend/base_backend.py | 9 ---- .../mode_backend/chunked_prefill/impl.py | 9 ++-- .../mode_backend/dp_backend/impl.py | 9 ++-- .../model_infer/mode_backend/eplb_manager.py | 29 +++++----- lightllm/utils/envs_utils.py | 8 +-- unit_tests/common/fused_moe/test_eplb.py | 53 ++++++------------- 6 files changed, 43 insertions(+), 74 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 6294beb155..4a72ed66ba 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -303,15 +303,6 @@ def infer_loop(self): def prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): raise NotImplementedError() - def _run_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): - self.prefill(event_pack=event_pack, prefill_reqs=prefill_reqs) - if self.eplb_manager is not None: - self.eplb_manager.step() - - def _poll_eplb(self): - if self.eplb_manager is not None: - self.eplb_manager.poll() - def decode(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): raise NotImplementedError() diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 085c66d618..28ef888b4e 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -61,11 +61,12 @@ def infer_loop(self): event_pack.wait_to_forward() - # Keep EPLB collectives ordered before normal collectives/forward. - self._poll_eplb() - self._try_read_new_reqs() + # Keep EPLB collectives ordered before normal collectives/forward. + if self.eplb_manager is not None: + self.eplb_manager.step() + prefill_reqs, decode_reqs = self._get_classed_reqs( no_decode=self.classed_req_no_decode, strict_prefill=self.classed_req_strict_prefill, @@ -78,7 +79,7 @@ def infer_loop(self): # 进行一次流同步,保证 _try_read_new_reqs 中的一些算子操作,必然已经完成。 # 防止后续的推理流程读取到显存中可能存在错误的数据。 g_infer_context.get_overlap_stream().wait_stream(torch.cuda.current_stream()) - self._run_prefill( + self.prefill( event_pack=event_pack, prefill_reqs=prefill_reqs, ) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 793b729f82..387da3d07f 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -122,11 +122,12 @@ def infer_loop(self): event_pack.wait_to_forward() - # Keep EPLB collectives ordered before normal collectives/forward. - self._poll_eplb() - self._try_read_new_reqs() + # Keep EPLB collectives ordered before normal collectives/forward. + if self.eplb_manager is not None: + self.eplb_manager.step() + prefill_reqs, decode_reqs = self._get_classed_reqs( no_decode=self.classed_req_no_decode, strict_prefill=self.classed_req_strict_prefill, @@ -148,7 +149,7 @@ def infer_loop(self): # 进行一次流同步,保证 _try_read_new_reqs 中的一些算子操作,必然已经完成。 # 防止后续的推理流程读取到显存中可能存在错误的数据。 g_infer_context.get_overlap_stream().wait_stream(torch.cuda.current_stream()) - self._run_prefill( + self.prefill( event_pack=event_pack, prefill_reqs=prefill_reqs, ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index a95a3f9c33..2b38f01738 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -30,7 +30,7 @@ ) from lightllm.utils.envs_utils import ( get_eplb_rebalance_gain_threshold, - get_prefill_eplb_step_interval, + get_eplb_step_interval, ) from lightllm.utils.log_utils import init_logger from lightllm.utils.shm_port_args import get_shm_port_args @@ -57,8 +57,8 @@ def __init__(self, model: TpPartBaseModel) -> None: self.num_logical_experts: int = routed.pop() self.num_redundant_experts_per_rank: int = redundant.pop() self.num_primary_experts_per_rank: int = self.num_logical_experts // self.world_size - self.step_interval: int = get_prefill_eplb_step_interval() - self.prefill_steps: int = 0 + self.step_interval: int = get_eplb_step_interval() + self.steps: int = 0 self.next_evaluation_step: int = self.step_interval initial_local_expert_ids = build_initial_local_expert_ids( @@ -101,18 +101,17 @@ def __init__(self, model: TpPartBaseModel) -> None: f"step_interval={self.step_interval} planner={type(self.planner).__name__}" ) - def poll(self) -> None: - """在所有 rank 顺序一致的推理边界推进评估或专家传输。""" + def step(self) -> None: + """在推理边界推进评估、专家传输或负载采样。""" if self.in_flight_transfers: self._poll_in_flight() - elif self._evaluation is not None and self._evaluation_ready_on_all_ranks(): - self._finish_evaluation() - - def step(self) -> None: - if self.in_flight_transfers or self._evaluation is not None: return - self.prefill_steps += 1 - if self.prefill_steps >= self.next_evaluation_step: + if self._evaluation is not None: + if self._evaluation_ready_on_all_ranks(): + self._finish_evaluation() + return + self.steps += 1 + if self.steps >= self.next_evaluation_step: self._start_evaluation() def _restart_collection(self) -> None: @@ -120,7 +119,7 @@ def _restart_collection(self) -> None: torch._foreach_zero_(counters) for impl in self._eplb_impls: impl.recording = True - self.next_evaluation_step = self.prefill_steps + self.step_interval + self.next_evaluation_step = self.steps + self.step_interval def _collect_local_load(self) -> torch.Tensor: counters = [impl.route_counter for impl in self._eplb_impls] @@ -262,10 +261,10 @@ def _start_rebalance(self, result: Dict[str, Any]) -> None: if self.global_rank == 0: changed_slots = len(self.in_flight_transfers) logger.info( - "eplb started prefill_steps=%s max_before=%.4f max_after=%.4f " + "eplb started steps=%s max_before=%.4f max_after=%.4f " "p95_before=%.4f p95_after=%.4f rebalance_gain=%.4f " "changed_layer_count=%s changed_slot_count=%s", - self.prefill_steps, + self.steps, result["before"]["max"], result["after"]["max"], result["before"]["p95"], diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index c301f01949..a647e344b6 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -97,11 +97,11 @@ def get_lightllm_websocket_max_message_size(): @lru_cache(maxsize=None) -def get_prefill_eplb_step_interval(): - """Return the number of prefill forwards between EPLB attempts.""" - interval = int(os.getenv("LIGHTLLM_PREFILL_EPLB_STEP_INTERVAL", 20)) +def get_eplb_step_interval(): + """Return the number of inference steps between EPLB attempts.""" + interval = int(os.getenv("LIGHTLLM_EPLB_STEP_INTERVAL", 20)) if interval <= 0: - raise ValueError("LIGHTLLM_PREFILL_EPLB_STEP_INTERVAL must be greater than 0") + raise ValueError("LIGHTLLM_EPLB_STEP_INTERVAL must be greater than 0") return interval diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 49967bc6fa..97a1f0305d 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -25,7 +25,6 @@ from lightllm.server.router.model_infer.mode_backend import ( eplb_transfer as transfer_module, ) -from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( deepgemm_impl as deepgemm_module, ) @@ -673,7 +672,7 @@ def test_steady_sampling_resets_aggregated_route_counter(): world_size=1, ) manager._eplb_impls = [impl] - manager.prefill_steps = 0 + manager.steps = 0 manager.step_interval = 20 manager._restart_collection() @@ -1417,25 +1416,25 @@ def is_finished(self): g_infer_context.overlap_stream = original_overlap_stream -def test_manager_inflight_step_does_not_poll(): +def test_manager_step_advances_inflight_transfer(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1, 2)] manager._evaluation = None calls = [] manager._poll_in_flight = lambda: calls.append("poll") manager.step() - assert calls == [] + assert calls == ["poll"] def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.in_flight_transfers = [] manager._evaluation = None - manager.prefill_steps = 0 + manager.steps = 0 manager.step_interval = 3 manager.next_evaluation_step = 3 starts = [] - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(manager.prefill_steps)) + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(manager.steps)) manager.step() manager.step() @@ -1447,7 +1446,7 @@ def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): def test_manager_restart_collection_resets_counters_and_arms_next_window(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.prefill_steps = 11 + manager.steps = 11 manager.step_interval = 20 counter = torch.ones(4, dtype=torch.int64) impl = _test_moe_impl( @@ -1466,7 +1465,7 @@ def test_manager_restart_collection_resets_counters_and_arms_next_window(): assert manager.next_evaluation_step == 31 -def test_manager_poll_advances_inflight_before_evaluation(): +def test_manager_step_advances_inflight_before_evaluation(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1, 2)] manager._evaluation = Future() @@ -1474,12 +1473,12 @@ def test_manager_poll_advances_inflight_before_evaluation(): manager._poll_in_flight = lambda: calls.append("inflight") manager._finish_evaluation = lambda: calls.append("evaluation") - manager.poll() + manager.step() assert calls == ["inflight"] -def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): +def test_manager_step_waits_for_all_evaluation_results(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.in_flight_transfers = [] manager._evaluation = Future() @@ -1489,40 +1488,18 @@ def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): manager._finish_evaluation = lambda: calls.append("evaluation") monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(0)) - manager.poll() + manager.step() assert calls == [] manager._evaluation.set_result({"kind": "no_improvement"}) monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) - manager.poll() + manager.step() assert calls == ["evaluation"] -def test_mode_backend_owns_eplb_poll_and_prefill_step(): - backend = object.__new__(ModeBackend) - calls = [] - backend.eplb_manager = SimpleNamespace( - poll=lambda: calls.append("poll"), - step=lambda: calls.append("step"), - ) - backend.prefill = lambda **_kwargs: calls.append("prefill") - - backend._poll_eplb() - backend._run_prefill(event_pack=object(), prefill_reqs=[]) - - assert calls == ["poll", "prefill", "step"] - - -def test_mode_backend_eplb_hooks_are_noops_when_disabled(): - backend = object.__new__(ModeBackend) - backend.eplb_manager = None - calls = [] - backend.prefill = lambda **_kwargs: calls.append("prefill") - - backend._poll_eplb() - backend._run_prefill(event_pack=object(), prefill_reqs=[]) - - assert calls == ["prefill"] +def test_manager_exposes_one_lifecycle_step_entrypoint(): + assert hasattr(manager_module.EPLBManager, "step") + assert not hasattr(manager_module.EPLBManager, "poll") def test_pinned_transfer_copies_source_row_and_sends_to_destination(monkeypatch): @@ -1704,7 +1681,7 @@ def test_manager_constructs_pinned_memory_transfer(monkeypatch): monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) - monkeypatch.setattr(manager_module, "get_prefill_eplb_step_interval", lambda: 20) + monkeypatch.setattr(manager_module, "get_eplb_step_interval", lambda: 20) monkeypatch.setattr(manager_module, "get_eplb_rebalance_gain_threshold", lambda: 0.07) def new_group(*args, **kwargs): From dd9ccb333daa4f4c3a571fd511bf4df5d027f659 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 08:20:23 +0000 Subject: [PATCH 30/72] refactor(eplb): drive manager with explicit state machine --- .../model_infer/mode_backend/eplb_manager.py | 221 ++++++++++++------ unit_tests/common/fused_moe/test_eplb.py | 168 +++++++++---- 2 files changed, 276 insertions(+), 113 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 2b38f01738..7062310094 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -1,4 +1,5 @@ from concurrent.futures import Future +from enum import Enum import threading import time from typing import Any, Dict, List, Optional, Tuple @@ -41,26 +42,48 @@ EPLB_EXPERT_IMBALANCE_RATIO_METRIC = "lightllm_eplb_topk_expert_imbalance_ratio" +class EPLBManagerState(Enum): + """EPLB 管理器在一次负载均衡循环中的阶段。""" + + COLLECTING = "collecting" + EVALUATING = "evaluating" + TRANSFERRING = "transferring" + + class EPLBManager: - """收集专家负载、调用布局规划器,并在安全边界发布迁移后的专家权重。""" + """由 :meth:`step` 驱动的 EPLB 状态机。 + + 状态循环如下: + + ``COLLECTING -> EVALUATING -> TRANSFERRING -> COLLECTING`` + + 当规划器认为无需调整布局时,``EVALUATING`` 会直接回到 + ``COLLECTING``。每次调用 :meth:`step` 最多推进一个状态,耗时的负载 + 评估和权重传输在后台执行,主推理线程只负责轮询和提交结果。 + """ def __init__(self, model: TpPartBaseModel) -> None: weights: List[FusedMoeWeight] = _find_fused_moe_weights(model) assert weights, "EPLB requires at least one EP MoE layer" + + # 模型与专家拓扑:初始化后保持不变。 self._weights: List[FusedMoeWeight] = weights self.global_rank: int = get_global_rank() self.world_size: int = get_global_world_size() + assert self.world_size > 1, "EPLB requires more than one rank" self._eplb_impls = [weight.fuse_moe_impl for weight in weights] - routed = {impl.n_routed_experts for impl in self._eplb_impls} - redundant = {impl.num_redundant_experts_per_rank for impl in self._eplb_impls} - assert len(routed) == len(redundant) == 1 - self.num_logical_experts: int = routed.pop() - self.num_redundant_experts_per_rank: int = redundant.pop() + + first_impl = self._eplb_impls[0] + self.num_logical_experts: int = first_impl.n_routed_experts + self.num_redundant_experts_per_rank: int = first_impl.num_redundant_experts_per_rank self.num_primary_experts_per_rank: int = self.num_logical_experts // self.world_size + + # 评估调度:steps 只在 COLLECTING 状态递增。 self.step_interval: int = get_eplb_step_interval() self.steps: int = 0 self.next_evaluation_step: int = self.step_interval + # 专家布局与规划器:current_placement 只记录各层的冗余专家槽位。 initial_local_expert_ids = build_initial_local_expert_ids( self.num_logical_experts, self.world_size, @@ -80,19 +103,29 @@ def __init__(self, model: TpPartBaseModel) -> None: rebalance_gain_threshold=get_eplb_rebalance_gain_threshold(), ) - self.in_flight_transfers: List[EPLBTransferInfo] = [] - self.completed_transfers: List[PinnedMemoryEPLBTransfer] = [] + # 状态机运行数据:_enter_collecting() 负责设置初始状态。 + self.state: EPLBManagerState + + # EVALUATING 状态持有的后台评估任务。 + self._evaluation: Optional[Future] = None + + # TRANSFERRING 状态持有的传输计划、活动任务及目标布局。 + self.pending_transfer_infos: List[EPLBTransferInfo] = [] + self.completed_layer_transfers: List[PinnedMemoryEPLBTransfer] = [] self.active_transfer: Optional[PinnedMemoryEPLBTransfer] = None self.target_placement: Optional[torch.Tensor] = None self.target_metadata: Optional[Dict[int, torch.Tensor]] = None - self._evaluation: Optional[Future] = None - self.metric_client: Optional[MetricClient] = None + # 分布式通信:评估、状态同步和权重传输使用独立的通信组。 self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") self._control_status_buffer: torch.Tensor = torch.empty(1, dtype=torch.int32) - self._restart_collection() + + # 可观测性资源按需创建,避免非 rank 0 初始化无用客户端。 + self.metric_client: Optional[MetricClient] = None + + self._enter_collecting() if self.global_rank == 0: logger.info( @@ -102,30 +135,47 @@ def __init__(self, model: TpPartBaseModel) -> None: ) def step(self) -> None: - """在推理边界推进评估、专家传输或负载采样。""" - if self.in_flight_transfers: - self._poll_in_flight() + """在一个安全的推理边界推进一次状态机。""" + if self.state is EPLBManagerState.COLLECTING: + self._step_collecting() + return + + if self.state is EPLBManagerState.EVALUATING: + self._step_evaluating() return - if self._evaluation is not None: - if self._evaluation_ready_on_all_ranks(): - self._finish_evaluation() + + if self.state is EPLBManagerState.TRANSFERRING: + self._step_transferring() return + + raise RuntimeError(f"unknown EPLB manager state: {self.state!r}") + + def _step_collecting(self) -> None: + """记录一个采样步,并在采样窗口结束后开始评估。""" self.steps += 1 if self.steps >= self.next_evaluation_step: - self._start_evaluation() + self._enter_evaluating() + + def _enter_collecting(self) -> None: + """清理上一轮状态,并开始一个完整的负载采样窗口。""" + self.state = EPLBManagerState.COLLECTING + self._evaluation = None + self.pending_transfer_infos = [] + self.completed_layer_transfers = [] + self.active_transfer = None + self.target_placement = None + self.target_metadata = None - def _restart_collection(self) -> None: - counters: List[torch.Tensor] = [impl.route_counter for impl in self._eplb_impls] - torch._foreach_zero_(counters) - for impl in self._eplb_impls: - impl.recording = True self.next_evaluation_step = self.steps + self.step_interval - def _collect_local_load(self) -> torch.Tensor: + def _snapshot_local_load(self) -> torch.Tensor: + """在当前计算流中复制计数快照,并立即开始下一个统计窗口。""" counters = [impl.route_counter for impl in self._eplb_impls] if any(counter.ndim != 1 or counter.shape[0] != self.num_logical_experts for counter in counters): raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") - return torch.stack(counters).cpu() + local_load = torch.stack(counters) + torch._foreach_zero_(counters) + return local_load def _plan_and_broadcast(self, global_load: torch.Tensor) -> Dict[str, Any]: result: Optional[Dict[str, Any]] = None @@ -139,10 +189,9 @@ def _plan_and_broadcast(self, global_load: torch.Tensor) -> Dict[str, Any]: except BaseException as exc: local_error = exc result = {"kind": "error", "message": f"{type(exc).__name__}: {exc}"} - if self.world_size > 1: - values = [result] - dist.broadcast_object_list(values, src=0, group=self.evaluation_group) - result = values[0] + values = [result] + dist.broadcast_object_list(values, src=0, group=self.evaluation_group) + result = values[0] if result["kind"] == "error": if local_error is not None: raise RuntimeError("EPLB planner failed on rank zero") from local_error @@ -194,11 +243,16 @@ def _build_rebalance_data(self, result: Dict[str, Any]) -> Tuple[Dict[int, torch ) return metadata_by_layer, planned_transfers - def _evaluate_after_event(self, event: torch.cuda.Event, evaluation: Future) -> None: + def _evaluate_after_event( + self, + event: torch.cuda.Event, + local_load: torch.Tensor, + evaluation: Future, + ) -> None: try: - torch.cuda.set_device(self._eplb_impls[0].route_counter.device) + torch.cuda.set_device(local_load.device) event.synchronize() - global_load = self._collect_local_load() + global_load = local_load.cpu() dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) result = self._plan_and_broadcast(global_load) result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) @@ -208,19 +262,22 @@ def _evaluate_after_event(self, event: torch.cuda.Event, evaluation: Future) -> except BaseException as exc: evaluation.set_exception(exc) - def _start_evaluation(self) -> None: - for impl in self._eplb_impls: - impl.recording = False + def _enter_evaluating(self) -> None: + """截取本轮计数,并在后台汇总负载和计算新布局。""" + local_load = self._snapshot_local_load() event = torch.cuda.Event() event.record(torch.cuda.current_stream()) - self._evaluation = Future() + evaluation = Future() + self._evaluation = evaluation + self.state = EPLBManagerState.EVALUATING threading.Thread( target=self._evaluate_after_event, - args=(event, self._evaluation), + args=(event, local_load, evaluation), daemon=True, ).start() def _evaluation_ready_on_all_ranks(self) -> bool: + assert self._evaluation is not None ready = self._evaluation.done() error = self._evaluation.exception() if ready else None status = EPLB_CONTROL_ERROR if error is not None else int(ready) @@ -233,16 +290,22 @@ def _evaluation_ready_on_all_ranks(self) -> bool: raise RuntimeError("EPLB evaluation failed on another rank") return bool(global_evaluation_status_value) - def _finish_evaluation(self) -> None: + def _step_evaluating(self) -> None: + """等待所有 rank 完成评估,然后进入采样或传输状态。""" + if not self._evaluation_ready_on_all_ranks(): + return + + assert self._evaluation is not None result = self._evaluation.result() self._evaluation = None self._publish_expert_load_metric(result) if result["kind"] != "planned": if self.global_rank == 0: logger.info("eplb skip rearrangement kind=%s", result["kind"]) - self._restart_collection() + self._enter_collecting() return - self._start_rebalance(result) + + self._enter_transferring(result) def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: if self.global_rank != 0: @@ -251,15 +314,19 @@ def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: self.metric_client = MetricClient(get_shm_port_args().metric_port) self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) - def _start_rebalance(self, result: Dict[str, Any]) -> None: + def _enter_transferring(self, result: Dict[str, Any]) -> None: + """安装传输计划,并启动计划中的第一个专家传输。""" self.target_placement = torch.tensor(result["placement"], dtype=torch.int64) self.target_metadata = result["metadata"] - self.in_flight_transfers = result["transfer_infos"] - self.completed_transfers = [] - self.in_flight_started_at = time.time() + self.pending_transfer_infos = list(result["transfer_infos"]) + if not self.pending_transfer_infos: + raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") + self.completed_layer_transfers = [] + self.rebalance_started_at = time.time() + self.state = EPLBManagerState.TRANSFERRING self._start_next_transfer() if self.global_rank == 0: - changed_slots = len(self.in_flight_transfers) + changed_slots = len(self.pending_transfer_infos) logger.info( "eplb started steps=%s max_before=%.4f max_after=%.4f " "p95_before=%.4f p95_after=%.4f rebalance_gain=%.4f " @@ -274,37 +341,58 @@ def _start_rebalance(self, result: Dict[str, Any]) -> None: changed_slots, ) - def _poll_in_flight(self) -> None: + def _step_transferring(self) -> None: + """推进当前传输,并在一层完成后原子地发布该层。""" + if not self._active_transfer_finished_on_all_ranks(): + return + + completed_info = self._complete_active_transfer() + if self._next_transfer_is_in_layer(completed_info.layer_index): + self._start_next_transfer() + return + + self._synchronize_and_commit_layer(completed_info.layer_index) + if self.pending_transfer_infos: + self._start_next_transfer() + return + + self.active_transfer = None + self._finish_rebalance() + + def _active_transfer_finished_on_all_ranks(self) -> bool: + """仅当所有 rank 都完成当前传输时返回 ``True``。""" assert self.active_transfer is not None all_ranks_finished = self._control_status_buffer.fill_(int(self.active_transfer.is_finished())) dist.all_reduce(all_ranks_finished, op=dist.ReduceOp.MIN, group=self.control_group) - if not bool(all_ranks_finished.item()): - return - expected_transfer_info: EPLBTransferInfo = self.in_flight_transfers[0] + return bool(all_ranks_finished.item()) + + def _complete_active_transfer(self) -> EPLBTransferInfo: + """将当前传输从待处理队列移动到本层的已完成列表。""" + assert self.active_transfer is not None + assert self.pending_transfer_infos + + expected_transfer_info: EPLBTransferInfo = self.pending_transfer_infos[0] if self.active_transfer.transfer_info != expected_transfer_info: raise RuntimeError("EPLB completed transfer does not match the expected transfer info") - self.completed_transfers.append(self.active_transfer) - self.in_flight_transfers.pop(0) + self.completed_layer_transfers.append(self.active_transfer) + self.pending_transfer_infos.pop(0) + return expected_transfer_info - if self.in_flight_transfers and self.in_flight_transfers[0].layer_index == expected_transfer_info.layer_index: - self._start_next_transfer() - return + def _next_transfer_is_in_layer(self, layer_index: int) -> bool: + """判断下一个待处理传输是否仍属于当前层。""" + return bool(self.pending_transfer_infos and self.pending_transfer_infos[0].layer_index == layer_index) + def _synchronize_and_commit_layer(self, layer_index: int) -> None: + """等待旧权重使用完毕,然后发布一层的新权重和 metadata。""" from lightllm.server.router.model_infer.infer_batch import g_infer_context torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - layer_index: int = expected_transfer_info.layer_index self._commit_transferred_layer(layer_index) - self.completed_transfers.clear() - if self.in_flight_transfers: - self._start_next_transfer() - else: - self.active_transfer = None - self._finish_rebalance() + self.completed_layer_transfers.clear() def _start_next_transfer(self) -> None: - transfer_info: EPLBTransferInfo = self.in_flight_transfers[0] + transfer_info: EPLBTransferInfo = self.pending_transfer_infos[0] self.active_transfer = PinnedMemoryEPLBTransfer( self._weights, self.transfer_group, @@ -317,7 +405,7 @@ def _commit_transferred_layer(self, layer_index: int) -> None: """在主推理线程中同步发布一层权重和路由 metadata。""" assert self.target_placement is not None target_redundant_expert_ids: List[int] = self.target_placement[layer_index, self.global_rank].tolist() - for transfer in self.completed_transfers: + for transfer in self.completed_layer_transfers: transfer_info: EPLBTransferInfo = transfer.transfer_info if transfer_info.dest_rank != self.global_rank: continue @@ -335,13 +423,12 @@ def _commit_layer_metadata(self, layer_index: int) -> None: def _finish_rebalance(self) -> None: assert self.target_placement is not None self.current_placement = self.target_placement - self.target_placement = None - self.target_metadata = None - self._restart_collection() + elapsed = time.time() - self.rebalance_started_at + self._enter_collecting() if self.global_rank == 0: logger.info( "eplb completed wall_time=%.2fs", - time.time() - self.in_flight_started_at, + elapsed, ) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 97a1f0305d..91262589a0 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -548,7 +548,7 @@ def test_transfer_plan_respects_explicit_target_slots(): } -def test_manager_collects_logical_route_counters(): +def test_manager_snapshots_and_resets_logical_route_counters(): counters = [ torch.tensor([10, 11], dtype=torch.int64), torch.tensor([40, 41], dtype=torch.int64), @@ -576,18 +576,19 @@ def test_manager_collects_logical_route_counters(): weight.layer_num_ = layer_num manager.num_logical_experts = 2 - samples = manager._collect_local_load() + samples = manager._snapshot_local_load() assert torch.equal( samples, torch.tensor([[10, 11], [40, 41]], dtype=torch.int64), ) + assert all(torch.count_nonzero(counter) == 0 for counter in counters) -def test_manager_delegates_distribution_planning_to_planner_class(): +def test_manager_delegates_distribution_planning_to_planner_class(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.global_rank = 0 - manager.world_size = 1 + manager.world_size = 2 manager.evaluation_group = object() manager.current_placement = torch.tensor([[[1]]]) logical_load = torch.tensor([[10, 20]]) @@ -598,6 +599,12 @@ def as_dict(self): return {"kind": "no_improvement"} manager.planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or Result())) + broadcasts = [] + monkeypatch.setattr( + manager_module.dist, + "broadcast_object_list", + lambda values, **kwargs: broadcasts.append((values, kwargs)), + ) result = manager._plan_and_broadcast(logical_load) @@ -605,6 +612,7 @@ def as_dict(self): assert len(calls) == 1 assert calls[0][0] == logical_load.tolist() assert calls[0][1] == manager.current_placement.tolist() + assert broadcasts == [([result], {"src": 0, "group": manager.evaluation_group})] def test_expert_load_imbalance_ratio_averages_layer_ratios(): @@ -661,13 +669,13 @@ def cpu_zeros(*shape, **kwargs): assert impl.route_counter.shape == (4,) -def test_steady_sampling_resets_aggregated_route_counter(): +def test_enter_collecting_keeps_counts_accumulated_during_evaluation(): counter = torch.ones((4,), dtype=torch.int64) manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) impl = _test_moe_impl( eplb=True, route_counter=counter, - recording=False, + recording=True, num_logical_experts=4, world_size=1, ) @@ -675,11 +683,11 @@ def test_steady_sampling_resets_aggregated_route_counter(): manager.steps = 0 manager.step_interval = 20 - manager._restart_collection() - manager._restart_collection() + manager._enter_collecting() assert impl.route_counter.shape == (4,) - assert torch.count_nonzero(impl.route_counter) == 0 + assert torch.equal(impl.route_counter, torch.ones(4, dtype=torch.int64)) + assert impl.recording def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): @@ -729,7 +737,6 @@ def test_manager_evaluation_collective_sums_rank_loads(monkeypatch): manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0) manager.evaluation_group = object() local = torch.full((1, 4), 100, dtype=torch.int64) - manager._collect_local_load = lambda: local.clone() seen = {} def all_reduce(tensor, **kwargs): @@ -748,7 +755,11 @@ def plan_and_broadcast(global_load): manager._plan_and_broadcast = plan_and_broadcast evaluation = Future() - manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})(), evaluation) + manager._evaluate_after_event( + type("Event", (), {"synchronize": lambda self: None})(), + local.clone(), + evaluation, + ) result = evaluation.result() assert seen["group"] is manager.evaluation_group @@ -788,7 +799,7 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c manager.num_redundant_experts_per_rank = 1 manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() manager.evaluation_group = object() - manager._collect_local_load = lambda: torch.full((3, 4), 100, dtype=torch.int64) + local = torch.full((3, 4), 100, dtype=torch.int64) planned_placement = torch.tensor( [ [[3], [0], [1], [2]], @@ -819,7 +830,11 @@ def build_maps_for_layers(*args, **kwargs): ) evaluation = Future() - manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})(), evaluation) + manager._evaluate_after_event( + type("Event", (), {"synchronize": lambda self: None})(), + local, + evaluation, + ) result = evaluation.result() assert calls == [(2, 4, 2)] @@ -1245,7 +1260,7 @@ def test_manager_commits_completed_transfer_rows_by_planned_destination_index(): manager.target_placement = torch.tensor([[[7, 6]]]) manager._eplb_impls = [SimpleNamespace(local_logics_expert_ids_list=local_expert_ids)] manager._commit_layer_metadata = lambda _layer: None - manager.completed_transfers = [ + manager.completed_layer_transfers = [ SimpleNamespace( transfer_info=EPLBTransferInfo(1, 0, 7, 0, 3), tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -7))], @@ -1280,14 +1295,14 @@ def is_finished(self): manager.active_transfer = Transfer(info0) manager.control_group = object() manager._control_status_buffer = torch.empty(1, dtype=torch.int32) - manager.in_flight_transfers = [info0, info0b, info1] - manager.completed_transfers = [] + manager.pending_transfer_infos = [info0, info0b, info1] + manager.completed_layer_transfers = [] committed, finished = [], [] manager._commit_transferred_layer = committed.append manager._finish_rebalance = lambda: finished.append(True) def start_next_transfer(): - manager.active_transfer = Transfer(manager.in_flight_transfers[0]) + manager.active_transfer = Transfer(manager.pending_transfer_infos[0]) manager._start_next_transfer = start_next_transfer operations = [] @@ -1304,28 +1319,28 @@ def set_global_ready(count): return lambda tensor, **_kwargs: tensor.fill_(count) monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(0)) - manager._poll_in_flight() + manager._step_transferring() assert committed == [] monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) - manager._poll_in_flight() + manager._step_transferring() assert committed == [] assert operations == [] assert not finished - manager._poll_in_flight() + manager._step_transferring() assert committed == [0] assert operations == [("wait", overlap_stream)] - manager._poll_in_flight() + manager._step_transferring() assert committed == [0, 1] assert finished == [True] expected = EPLBTransferInfo(0, 2, 5, 1, 2) - manager.in_flight_transfers = [expected] + manager.pending_transfer_infos = [expected] manager.active_transfer = Transfer(EPLBTransferInfo(1, 2, 5, 1, 2)) with pytest.raises(RuntimeError, match="does not match"): - manager._poll_in_flight() + manager._step_transferring() def test_evaluation_ready_gate_propagates_local_and_remote_errors(monkeypatch): @@ -1387,8 +1402,8 @@ def is_finished(self): manager.active_transfer = Transfer(live, received, transfer_info) manager.control_group = object() manager._control_status_buffer = torch.empty(1, dtype=torch.int32) - manager.in_flight_transfers = [transfer_info] - manager.completed_transfers = [] + manager.pending_transfer_infos = [transfer_info] + manager.completed_layer_transfers = [] manager.num_primary_experts_per_rank = 0 manager.global_rank = 0 manager.target_placement = torch.tensor([[[2]]]) @@ -1404,7 +1419,7 @@ def is_finished(self): torch.cuda._sleep(20_000_000) previous_read.copy_(live, non_blocking=True) with torch.cuda.stream(destination_stream): - manager._poll_in_flight() + manager._step_transferring() with torch.cuda.stream(source_stream): source_stream.wait_stream(destination_stream) next_read.copy_(live, non_blocking=True) @@ -1418,23 +1433,23 @@ def is_finished(self): def test_manager_step_advances_inflight_transfer(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1, 2)] - manager._evaluation = None + manager.state = manager_module.EPLBManagerState.TRANSFERRING calls = [] - manager._poll_in_flight = lambda: calls.append("poll") + manager._step_transferring = lambda: calls.append("transfer") + manager.step() - assert calls == ["poll"] + + assert calls == ["transfer"] def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_transfers = [] - manager._evaluation = None + manager.state = manager_module.EPLBManagerState.COLLECTING manager.steps = 0 manager.step_interval = 3 manager.next_evaluation_step = 3 starts = [] - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(manager.steps)) + monkeypatch.setattr(manager, "_enter_evaluating", lambda: starts.append(manager.steps)) manager.step() manager.step() @@ -1444,7 +1459,32 @@ def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): assert starts == [3] -def test_manager_restart_collection_resets_counters_and_arms_next_window(): +def test_manager_enters_evaluating_without_stopping_route_counting(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + impl = SimpleNamespace(recording=True) + local_load = torch.tensor([[1, 2]], dtype=torch.int64) + event = SimpleNamespace(record=lambda stream: None) + worker = SimpleNamespace(start=lambda: None) + thread_args = [] + manager._eplb_impls = [impl] + manager._snapshot_local_load = lambda: local_load + + monkeypatch.setattr(manager_module.torch.cuda, "Event", lambda: event) + monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: object()) + monkeypatch.setattr( + manager_module.threading, + "Thread", + lambda **kwargs: (thread_args.append(kwargs["args"]) or worker), + ) + + manager._enter_evaluating() + + assert impl.recording + assert manager.state is manager_module.EPLBManagerState.EVALUATING + assert thread_args[0][1] is local_load + + +def test_manager_enter_collecting_clears_state_and_arms_next_window(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.steps = 11 manager.step_interval = 20 @@ -1452,40 +1492,65 @@ def test_manager_restart_collection_resets_counters_and_arms_next_window(): impl = _test_moe_impl( eplb=True, route_counter=counter, - recording=False, + recording=True, num_logical_experts=4, world_size=1, ) manager._eplb_impls = [impl] - manager._restart_collection() + manager._enter_collecting() - assert torch.count_nonzero(counter) == 0 + assert torch.equal(counter, torch.ones(4, dtype=torch.int64)) assert impl.recording + assert manager.state is manager_module.EPLBManagerState.COLLECTING assert manager.next_evaluation_step == 31 -def test_manager_step_advances_inflight_before_evaluation(): +def test_manager_step_uses_explicit_state_instead_of_pending_work(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_transfers = [EPLBTransferInfo(0, 0, 2, 1, 2)] + manager.state = manager_module.EPLBManagerState.TRANSFERRING + manager.pending_transfer_infos = [EPLBTransferInfo(0, 0, 2, 1, 2)] manager._evaluation = Future() calls = [] - manager._poll_in_flight = lambda: calls.append("inflight") - manager._finish_evaluation = lambda: calls.append("evaluation") + manager._step_transferring = lambda: calls.append("transfer") + manager._step_evaluating = lambda: calls.append("evaluation") manager.step() - assert calls == ["inflight"] + assert calls == ["transfer"] + + +def test_manager_enters_transferring_state_with_planned_work(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) + starts = [] + manager.global_rank = 1 + manager._start_next_transfer = lambda: starts.append(manager.pending_transfer_infos[0]) + + manager._enter_transferring( + { + "placement": [[[2], [3]]], + "metadata": {0: torch.tensor([1])}, + "transfer_infos": [transfer_info], + } + ) + + assert manager.state is manager_module.EPLBManagerState.TRANSFERRING + assert manager.pending_transfer_infos == [transfer_info] + assert manager.completed_layer_transfers == [] + assert starts == [transfer_info] def test_manager_step_waits_for_all_evaluation_results(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight_transfers = [] + manager.state = manager_module.EPLBManagerState.EVALUATING manager._evaluation = Future() manager.control_group = object() manager._control_status_buffer = torch.empty(1, dtype=torch.int32) calls = [] - manager._finish_evaluation = lambda: calls.append("evaluation") + manager._publish_expert_load_metric = lambda _result: calls.append("publish") + manager._enter_collecting = lambda: calls.append("collect") + manager.global_rank = 1 monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(0)) manager.step() @@ -1494,7 +1559,7 @@ def test_manager_step_waits_for_all_evaluation_results(monkeypatch): manager._evaluation.set_result({"kind": "no_improvement"}) monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) manager.step() - assert calls == ["evaluation"] + assert calls == ["publish", "collect"] def test_manager_exposes_one_lifecycle_step_entrypoint(): @@ -1658,6 +1723,15 @@ def synchronize(self): transfer.start() +def test_manager_requires_more_than_one_rank(monkeypatch): + monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [object()]) + monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 1) + + with pytest.raises(AssertionError, match="more than one rank"): + manager_module.EPLBManager(type("Model", (), {})()) + + def test_manager_constructs_pinned_memory_transfer(monkeypatch): weight = type( "Weight", @@ -1667,6 +1741,7 @@ def test_manager_constructs_pinned_memory_transfer(monkeypatch): "layer_num_": 0, "fuse_moe_impl": _test_moe_impl( eplb=True, + recording=True, num_logical_experts=4, world_size=2, num_redundant_experts_per_rank=2, @@ -1699,6 +1774,7 @@ def new_group(*args, **kwargs): monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) manager = manager_module.EPLBManager(type("Model", (), {})()) assert manager.active_transfer is None + assert manager.state is manager_module.EPLBManagerState.COLLECTING assert ( manager.evaluation_group, manager.control_group, @@ -1706,7 +1782,7 @@ def new_group(*args, **kwargs): ) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) - manager.in_flight_transfers = [transfer_info] + manager.pending_transfer_infos = [transfer_info] manager._start_next_transfer() assert transfer_calls == [([weight], groups[2], 0, transfer_info)] assert transfer_starts == [True] From 2a6689d736c6071442365213a96da66911a249c1 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 08:24:53 +0000 Subject: [PATCH 31/72] refactor(eplb): scope runtime state to lifecycle phases --- .../model_infer/mode_backend/eplb_manager.py | 61 ++++++------------- unit_tests/common/fused_moe/test_eplb.py | 37 ++++++++++- 2 files changed, 55 insertions(+), 43 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 7062310094..25f6139127 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -81,7 +81,6 @@ def __init__(self, model: TpPartBaseModel) -> None: # 评估调度:steps 只在 COLLECTING 状态递增。 self.step_interval: int = get_eplb_step_interval() self.steps: int = 0 - self.next_evaluation_step: int = self.step_interval # 专家布局与规划器:current_placement 只记录各层的冗余专家槽位。 initial_local_expert_ids = build_initial_local_expert_ids( @@ -103,28 +102,12 @@ def __init__(self, model: TpPartBaseModel) -> None: rebalance_gain_threshold=get_eplb_rebalance_gain_threshold(), ) - # 状态机运行数据:_enter_collecting() 负责设置初始状态。 - self.state: EPLBManagerState - - # EVALUATING 状态持有的后台评估任务。 - self._evaluation: Optional[Future] = None - - # TRANSFERRING 状态持有的传输计划、活动任务及目标布局。 - self.pending_transfer_infos: List[EPLBTransferInfo] = [] - self.completed_layer_transfers: List[PinnedMemoryEPLBTransfer] = [] - self.active_transfer: Optional[PinnedMemoryEPLBTransfer] = None - self.target_placement: Optional[torch.Tensor] = None - self.target_metadata: Optional[Dict[int, torch.Tensor]] = None - # 分布式通信:评估、状态同步和权重传输使用独立的通信组。 self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") self._control_status_buffer: torch.Tensor = torch.empty(1, dtype=torch.int32) - # 可观测性资源按需创建,避免非 rank 0 初始化无用客户端。 - self.metric_client: Optional[MetricClient] = None - self._enter_collecting() if self.global_rank == 0: @@ -157,15 +140,8 @@ def _step_collecting(self) -> None: self._enter_evaluating() def _enter_collecting(self) -> None: - """清理上一轮状态,并开始一个完整的负载采样窗口。""" + """开始一个完整的负载采样窗口。""" self.state = EPLBManagerState.COLLECTING - self._evaluation = None - self.pending_transfer_infos = [] - self.completed_layer_transfers = [] - self.active_transfer = None - self.target_placement = None - self.target_metadata = None - self.next_evaluation_step = self.steps + self.step_interval def _snapshot_local_load(self) -> torch.Tensor: @@ -268,6 +244,7 @@ def _enter_evaluating(self) -> None: event = torch.cuda.Event() event.record(torch.cuda.current_stream()) evaluation = Future() + del self.next_evaluation_step self._evaluation = evaluation self.state = EPLBManagerState.EVALUATING threading.Thread( @@ -277,7 +254,6 @@ def _enter_evaluating(self) -> None: ).start() def _evaluation_ready_on_all_ranks(self) -> bool: - assert self._evaluation is not None ready = self._evaluation.done() error = self._evaluation.exception() if ready else None status = EPLB_CONTROL_ERROR if error is not None else int(ready) @@ -295,9 +271,8 @@ def _step_evaluating(self) -> None: if not self._evaluation_ready_on_all_ranks(): return - assert self._evaluation is not None result = self._evaluation.result() - self._evaluation = None + del self._evaluation self._publish_expert_load_metric(result) if result["kind"] != "planned": if self.global_rank == 0: @@ -310,18 +285,22 @@ def _step_evaluating(self) -> None: def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: if self.global_rank != 0: return - if self.metric_client is None: - self.metric_client = MetricClient(get_shm_port_args().metric_port) - self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) + metric_client = getattr(self, "metric_client", None) + if metric_client is None: + metric_client = MetricClient(get_shm_port_args().metric_port) + self.metric_client = metric_client + metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) def _enter_transferring(self, result: Dict[str, Any]) -> None: """安装传输计划,并启动计划中的第一个专家传输。""" + pending_transfer_infos = list(result["transfer_infos"]) + if not pending_transfer_infos: + raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") + self.target_placement = torch.tensor(result["placement"], dtype=torch.int64) self.target_metadata = result["metadata"] - self.pending_transfer_infos = list(result["transfer_infos"]) - if not self.pending_transfer_infos: - raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") - self.completed_layer_transfers = [] + self.pending_transfer_infos = pending_transfer_infos + self.completed_layer_transfers: List[PinnedMemoryEPLBTransfer] = [] self.rebalance_started_at = time.time() self.state = EPLBManagerState.TRANSFERRING self._start_next_transfer() @@ -356,19 +335,16 @@ def _step_transferring(self) -> None: self._start_next_transfer() return - self.active_transfer = None self._finish_rebalance() def _active_transfer_finished_on_all_ranks(self) -> bool: """仅当所有 rank 都完成当前传输时返回 ``True``。""" - assert self.active_transfer is not None all_ranks_finished = self._control_status_buffer.fill_(int(self.active_transfer.is_finished())) dist.all_reduce(all_ranks_finished, op=dist.ReduceOp.MIN, group=self.control_group) return bool(all_ranks_finished.item()) def _complete_active_transfer(self) -> EPLBTransferInfo: """将当前传输从待处理队列移动到本层的已完成列表。""" - assert self.active_transfer is not None assert self.pending_transfer_infos expected_transfer_info: EPLBTransferInfo = self.pending_transfer_infos[0] @@ -403,7 +379,6 @@ def _start_next_transfer(self) -> None: def _commit_transferred_layer(self, layer_index: int) -> None: """在主推理线程中同步发布一层权重和路由 metadata。""" - assert self.target_placement is not None target_redundant_expert_ids: List[int] = self.target_placement[layer_index, self.global_rank].tolist() for transfer in self.completed_layer_transfers: transfer_info: EPLBTransferInfo = transfer.transfer_info @@ -417,13 +392,17 @@ def _commit_transferred_layer(self, layer_index: int) -> None: self._commit_layer_metadata(layer_index) def _commit_layer_metadata(self, layer_index: int) -> None: - assert self.target_metadata is not None self._eplb_impls[layer_index].logical_to_physical_map.copy_(self.target_metadata[layer_index]) def _finish_rebalance(self) -> None: - assert self.target_placement is not None self.current_placement = self.target_placement elapsed = time.time() - self.rebalance_started_at + del self.pending_transfer_infos + del self.completed_layer_transfers + del self.active_transfer + del self.target_placement + del self.target_metadata + del self.rebalance_started_at self._enter_collecting() if self.global_rank == 0: logger.info( diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 91262589a0..7ae96b8b7a 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1468,6 +1468,7 @@ def test_manager_enters_evaluating_without_stopping_route_counting(monkeypatch): thread_args = [] manager._eplb_impls = [impl] manager._snapshot_local_load = lambda: local_load + manager.next_evaluation_step = 20 monkeypatch.setattr(manager_module.torch.cuda, "Event", lambda: event) monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: object()) @@ -1482,9 +1483,10 @@ def test_manager_enters_evaluating_without_stopping_route_counting(monkeypatch): assert impl.recording assert manager.state is manager_module.EPLBManagerState.EVALUATING assert thread_args[0][1] is local_load + assert not hasattr(manager, "next_evaluation_step") -def test_manager_enter_collecting_clears_state_and_arms_next_window(): +def test_manager_enter_collecting_sets_state_and_arms_next_window(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.steps = 11 manager.step_interval = 20 @@ -1560,6 +1562,35 @@ def test_manager_step_waits_for_all_evaluation_results(monkeypatch): monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) manager.step() assert calls == ["publish", "collect"] + assert not hasattr(manager, "_evaluation") + + +def test_manager_finish_rebalance_releases_transferring_state(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + target_placement = torch.tensor([[[2], [3]]]) + manager.global_rank = 1 + manager.steps = 10 + manager.step_interval = 20 + manager.pending_transfer_infos = [] + manager.completed_layer_transfers = [] + manager.active_transfer = object() + manager.target_placement = target_placement + manager.target_metadata = {} + manager.rebalance_started_at = time.time() + + manager._finish_rebalance() + + assert manager.current_placement is target_placement + assert manager.state is manager_module.EPLBManagerState.COLLECTING + for attribute in ( + "pending_transfer_infos", + "completed_layer_transfers", + "active_transfer", + "target_placement", + "target_metadata", + "rebalance_started_at", + ): + assert not hasattr(manager, attribute) def test_manager_exposes_one_lifecycle_step_entrypoint(): @@ -1773,7 +1804,9 @@ def new_group(*args, **kwargs): logs = [] monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) manager = manager_module.EPLBManager(type("Model", (), {})()) - assert manager.active_transfer is None + assert not hasattr(manager, "_evaluation") + assert not hasattr(manager, "active_transfer") + assert not hasattr(manager, "target_placement") assert manager.state is manager_module.EPLBManagerState.COLLECTING assert ( manager.evaluation_group, From a3e055f3a38e4d66157fef6427110413dc011543 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 08:36:42 +0000 Subject: [PATCH 32/72] refactor(eplb): keep state transitions in step handlers --- .../model_infer/mode_backend/eplb_manager.py | 287 +++++++++--------- unit_tests/common/fused_moe/test_eplb.py | 126 ++++---- 2 files changed, 192 insertions(+), 221 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 25f6139127..75d022f82b 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -108,9 +108,11 @@ def __init__(self, model: TpPartBaseModel) -> None: self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") self._control_status_buffer: torch.Tensor = torch.empty(1, dtype=torch.int32) - self._enter_collecting() + self.state = EPLBManagerState.COLLECTING + self.next_evaluation_step = self.step_interval if self.global_rank == 0: + self.metric_client: MetricClient = MetricClient(get_shm_port_args().metric_port) logger.info( f"eplb enabled layers={len(weights)} num_logical_experts={self.num_logical_experts} " f"num_redundant_experts_per_rank={self.num_redundant_experts_per_rank} " @@ -133,16 +135,111 @@ def step(self) -> None: raise RuntimeError(f"unknown EPLB manager state: {self.state!r}") + # 状态处理:与 step() 的分发顺序保持一致。 + def _step_collecting(self) -> None: """记录一个采样步,并在采样窗口结束后开始评估。""" self.steps += 1 - if self.steps >= self.next_evaluation_step: - self._enter_evaluating() + if self.steps < self.next_evaluation_step: + return + + local_load = self._snapshot_local_load() + event = torch.cuda.Event() + event.record(torch.cuda.current_stream()) + evaluation = Future() + del self.next_evaluation_step + self._evaluation = evaluation + self.state = EPLBManagerState.EVALUATING + threading.Thread( + target=self._evaluate_after_event, + args=(event, local_load, evaluation), + daemon=True, + ).start() + + def _step_evaluating(self) -> None: + """等待所有 rank 完成评估,然后进入采样或传输状态。""" + if not self._evaluation_ready_on_all_ranks(): + return + + result = self._evaluation.result() + del self._evaluation + self._publish_expert_load_metric(result) + if result["kind"] != "planned": + if self.global_rank == 0: + logger.info("eplb skip rearrangement kind=%s", result["kind"]) + self.state = EPLBManagerState.COLLECTING + self.next_evaluation_step = self.steps + self.step_interval + return + + pending_transfer_infos = list(result["transfer_infos"]) + if not pending_transfer_infos: + raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") + + self.target_placement = torch.tensor(result["placement"], dtype=torch.int64) + self.target_metadata = result["metadata"] + self.pending_transfer_infos = pending_transfer_infos + self.completed_layer_transfers: List[PinnedMemoryEPLBTransfer] = [] + self.rebalance_started_at = time.time() + self.state = EPLBManagerState.TRANSFERRING + self._start_next_transfer() + if self.global_rank == 0: + logger.info( + "eplb started steps=%s max_before=%.4f max_after=%.4f " + "p95_before=%.4f p95_after=%.4f rebalance_gain=%.4f " + "changed_layer_count=%s changed_slot_count=%s", + self.steps, + result["before"]["max"], + result["after"]["max"], + result["before"]["p95"], + result["after"]["p95"], + result["rebalance_gain"], + result["changed_layer_count"], + len(self.pending_transfer_infos), + ) + + def _step_transferring(self) -> None: + """推进当前传输,并在一层完成后原子地发布该层。""" + if not self._active_transfer_finished_on_all_ranks(): + return + + completed_info = self._complete_active_transfer() + if self._next_transfer_is_in_layer(completed_info.layer_index): + self._start_next_transfer() + return + + self._synchronize_and_commit_layer(completed_info.layer_index) + if self.pending_transfer_infos: + self._start_next_transfer() + return - def _enter_collecting(self) -> None: - """开始一个完整的负载采样窗口。""" + elapsed = self._complete_rebalance() self.state = EPLBManagerState.COLLECTING self.next_evaluation_step = self.steps + self.step_interval + if self.global_rank == 0: + logger.info( + "eplb completed wall_time=%.2fs", + elapsed, + ) + + # 评估阶段内部实现。 + + def _evaluation_ready_on_all_ranks(self) -> bool: + ready = self._evaluation.done() + error = self._evaluation.exception() if ready else None + status = EPLB_CONTROL_ERROR if error is not None else int(ready) + global_evaluation_status = self._control_status_buffer.fill_(status) + dist.all_reduce(global_evaluation_status, op=dist.ReduceOp.MIN, group=self.control_group) + global_evaluation_status_value = int(global_evaluation_status.item()) + if global_evaluation_status_value < 0: + if error is not None: + raise RuntimeError("EPLB evaluation failed on this rank") from error + raise RuntimeError("EPLB evaluation failed on another rank") + return bool(global_evaluation_status_value) + + def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: + if self.global_rank != 0: + return + self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) def _snapshot_local_load(self) -> torch.Tensor: """在当前计算流中复制计数快照,并立即开始下一个统计窗口。""" @@ -153,6 +250,25 @@ def _snapshot_local_load(self) -> torch.Tensor: torch._foreach_zero_(counters) return local_load + def _evaluate_after_event( + self, + event: torch.cuda.Event, + local_load: torch.Tensor, + evaluation: Future, + ) -> None: + try: + torch.cuda.set_device(local_load.device) + event.synchronize() + global_load = local_load.cpu() + dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) + result = self._plan_and_broadcast(global_load) + result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) + if result["kind"] == "planned": + result["metadata"], result["transfer_infos"] = self._build_rebalance_data(result) + evaluation.set_result(result) + except BaseException as exc: + evaluation.set_exception(exc) + def _plan_and_broadcast(self, global_load: torch.Tensor) -> Dict[str, Any]: result: Optional[Dict[str, Any]] = None local_error: Optional[BaseException] = None @@ -219,123 +335,7 @@ def _build_rebalance_data(self, result: Dict[str, Any]) -> Tuple[Dict[int, torch ) return metadata_by_layer, planned_transfers - def _evaluate_after_event( - self, - event: torch.cuda.Event, - local_load: torch.Tensor, - evaluation: Future, - ) -> None: - try: - torch.cuda.set_device(local_load.device) - event.synchronize() - global_load = local_load.cpu() - dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) - result = self._plan_and_broadcast(global_load) - result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) - if result["kind"] == "planned": - result["metadata"], result["transfer_infos"] = self._build_rebalance_data(result) - evaluation.set_result(result) - except BaseException as exc: - evaluation.set_exception(exc) - - def _enter_evaluating(self) -> None: - """截取本轮计数,并在后台汇总负载和计算新布局。""" - local_load = self._snapshot_local_load() - event = torch.cuda.Event() - event.record(torch.cuda.current_stream()) - evaluation = Future() - del self.next_evaluation_step - self._evaluation = evaluation - self.state = EPLBManagerState.EVALUATING - threading.Thread( - target=self._evaluate_after_event, - args=(event, local_load, evaluation), - daemon=True, - ).start() - - def _evaluation_ready_on_all_ranks(self) -> bool: - ready = self._evaluation.done() - error = self._evaluation.exception() if ready else None - status = EPLB_CONTROL_ERROR if error is not None else int(ready) - global_evaluation_status = self._control_status_buffer.fill_(status) - dist.all_reduce(global_evaluation_status, op=dist.ReduceOp.MIN, group=self.control_group) - global_evaluation_status_value = int(global_evaluation_status.item()) - if global_evaluation_status_value < 0: - if error is not None: - raise RuntimeError("EPLB evaluation failed on this rank") from error - raise RuntimeError("EPLB evaluation failed on another rank") - return bool(global_evaluation_status_value) - - def _step_evaluating(self) -> None: - """等待所有 rank 完成评估,然后进入采样或传输状态。""" - if not self._evaluation_ready_on_all_ranks(): - return - - result = self._evaluation.result() - del self._evaluation - self._publish_expert_load_metric(result) - if result["kind"] != "planned": - if self.global_rank == 0: - logger.info("eplb skip rearrangement kind=%s", result["kind"]) - self._enter_collecting() - return - - self._enter_transferring(result) - - def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: - if self.global_rank != 0: - return - metric_client = getattr(self, "metric_client", None) - if metric_client is None: - metric_client = MetricClient(get_shm_port_args().metric_port) - self.metric_client = metric_client - metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) - - def _enter_transferring(self, result: Dict[str, Any]) -> None: - """安装传输计划,并启动计划中的第一个专家传输。""" - pending_transfer_infos = list(result["transfer_infos"]) - if not pending_transfer_infos: - raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") - - self.target_placement = torch.tensor(result["placement"], dtype=torch.int64) - self.target_metadata = result["metadata"] - self.pending_transfer_infos = pending_transfer_infos - self.completed_layer_transfers: List[PinnedMemoryEPLBTransfer] = [] - self.rebalance_started_at = time.time() - self.state = EPLBManagerState.TRANSFERRING - self._start_next_transfer() - if self.global_rank == 0: - changed_slots = len(self.pending_transfer_infos) - logger.info( - "eplb started steps=%s max_before=%.4f max_after=%.4f " - "p95_before=%.4f p95_after=%.4f rebalance_gain=%.4f " - "changed_layer_count=%s changed_slot_count=%s", - self.steps, - result["before"]["max"], - result["after"]["max"], - result["before"]["p95"], - result["after"]["p95"], - result["rebalance_gain"], - result["changed_layer_count"], - changed_slots, - ) - - def _step_transferring(self) -> None: - """推进当前传输,并在一层完成后原子地发布该层。""" - if not self._active_transfer_finished_on_all_ranks(): - return - - completed_info = self._complete_active_transfer() - if self._next_transfer_is_in_layer(completed_info.layer_index): - self._start_next_transfer() - return - - self._synchronize_and_commit_layer(completed_info.layer_index) - if self.pending_transfer_infos: - self._start_next_transfer() - return - - self._finish_rebalance() + # 传输阶段内部实现。 def _active_transfer_finished_on_all_ranks(self) -> bool: """仅当所有 rank 都完成当前传输时返回 ``True``。""" @@ -359,14 +359,6 @@ def _next_transfer_is_in_layer(self, layer_index: int) -> bool: """判断下一个待处理传输是否仍属于当前层。""" return bool(self.pending_transfer_infos and self.pending_transfer_infos[0].layer_index == layer_index) - def _synchronize_and_commit_layer(self, layer_index: int) -> None: - """等待旧权重使用完毕,然后发布一层的新权重和 metadata。""" - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - self._commit_transferred_layer(layer_index) - self.completed_layer_transfers.clear() - def _start_next_transfer(self) -> None: transfer_info: EPLBTransferInfo = self.pending_transfer_infos[0] self.active_transfer = PinnedMemoryEPLBTransfer( @@ -377,6 +369,25 @@ def _start_next_transfer(self) -> None: ) self.active_transfer.start() + def _synchronize_and_commit_layer(self, layer_index: int) -> None: + """等待旧权重使用完毕,然后发布一层的新权重和 metadata。""" + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) + self._commit_transferred_layer(layer_index) + self.completed_layer_transfers.clear() + + def _complete_rebalance(self) -> float: + self.current_placement = self.target_placement + elapsed = time.time() - self.rebalance_started_at + del self.pending_transfer_infos + del self.completed_layer_transfers + del self.active_transfer + del self.target_placement + del self.target_metadata + del self.rebalance_started_at + return elapsed + def _commit_transferred_layer(self, layer_index: int) -> None: """在主推理线程中同步发布一层权重和路由 metadata。""" target_redundant_expert_ids: List[int] = self.target_placement[layer_index, self.global_rank].tolist() @@ -394,22 +405,6 @@ def _commit_transferred_layer(self, layer_index: int) -> None: def _commit_layer_metadata(self, layer_index: int) -> None: self._eplb_impls[layer_index].logical_to_physical_map.copy_(self.target_metadata[layer_index]) - def _finish_rebalance(self) -> None: - self.current_placement = self.target_placement - elapsed = time.time() - self.rebalance_started_at - del self.pending_transfer_infos - del self.completed_layer_transfers - del self.active_transfer - del self.target_placement - del self.target_metadata - del self.rebalance_started_at - self._enter_collecting() - if self.global_rank == 0: - logger.info( - "eplb completed wall_time=%.2fs", - elapsed, - ) - def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: """Average each layer's maximum-to-mean logical-expert token ratio.""" diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 7ae96b8b7a..69b46d4227 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -646,6 +646,15 @@ def test_manager_publishes_expert_load_metrics_from_rank_zero(): ] +def test_manager_does_not_publish_expert_load_metrics_from_other_ranks(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.global_rank = 1 + + manager._publish_expert_load_metric({"expert_imbalance_ratio": 1.25}) + + assert not hasattr(manager, "metric_client") + + def test_eplb_route_counter_has_one_entry_per_logical_expert(monkeypatch): args = type( "Args", @@ -669,27 +678,6 @@ def cpu_zeros(*shape, **kwargs): assert impl.route_counter.shape == (4,) -def test_enter_collecting_keeps_counts_accumulated_during_evaluation(): - counter = torch.ones((4,), dtype=torch.int64) - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - impl = _test_moe_impl( - eplb=True, - route_counter=counter, - recording=True, - num_logical_experts=4, - world_size=1, - ) - manager._eplb_impls = [impl] - manager.steps = 0 - manager.step_interval = 20 - - manager._enter_collecting() - - assert impl.route_counter.shape == (4,) - assert torch.equal(impl.route_counter, torch.ones(4, dtype=torch.int64)) - assert impl.recording - - def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): args = type( "Args", @@ -1297,9 +1285,12 @@ def is_finished(self): manager._control_status_buffer = torch.empty(1, dtype=torch.int32) manager.pending_transfer_infos = [info0, info0b, info1] manager.completed_layer_transfers = [] + manager.global_rank = 1 + manager.steps = 0 + manager.step_interval = 20 committed, finished = [], [] manager._commit_transferred_layer = committed.append - manager._finish_rebalance = lambda: finished.append(True) + manager._complete_rebalance = lambda: (finished.append(True) or 0.0) def start_next_transfer(): manager.active_transfer = Transfer(manager.pending_transfer_infos[0]) @@ -1409,7 +1400,9 @@ def is_finished(self): manager.target_placement = torch.tensor([[[2]]]) manager._eplb_impls = [SimpleNamespace(local_logics_expert_ids_list=[1])] manager._commit_layer_metadata = lambda _layer: None - manager._finish_rebalance = lambda: None + manager._complete_rebalance = lambda: 0.0 + manager.steps = 0 + manager.step_interval = 20 monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) try: @@ -1444,31 +1437,15 @@ def test_manager_step_advances_inflight_transfer(): def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.state = manager_module.EPLBManagerState.COLLECTING - manager.steps = 0 - manager.step_interval = 3 - manager.next_evaluation_step = 3 - starts = [] - monkeypatch.setattr(manager, "_enter_evaluating", lambda: starts.append(manager.steps)) - - manager.step() - manager.step() - assert starts == [] - manager.step() - - assert starts == [3] - - -def test_manager_enters_evaluating_without_stopping_route_counting(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - impl = SimpleNamespace(recording=True) local_load = torch.tensor([[1, 2]], dtype=torch.int64) event = SimpleNamespace(record=lambda stream: None) worker = SimpleNamespace(start=lambda: None) thread_args = [] - manager._eplb_impls = [impl] + manager.state = manager_module.EPLBManagerState.COLLECTING + manager.steps = 0 + manager.step_interval = 3 + manager.next_evaluation_step = 3 manager._snapshot_local_load = lambda: local_load - manager.next_evaluation_step = 20 monkeypatch.setattr(manager_module.torch.cuda, "Event", lambda: event) monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: object()) @@ -1478,36 +1455,16 @@ def test_manager_enters_evaluating_without_stopping_route_counting(monkeypatch): lambda **kwargs: (thread_args.append(kwargs["args"]) or worker), ) - manager._enter_evaluating() + manager.step() + manager.step() + assert manager.state is manager_module.EPLBManagerState.COLLECTING + manager.step() - assert impl.recording assert manager.state is manager_module.EPLBManagerState.EVALUATING assert thread_args[0][1] is local_load assert not hasattr(manager, "next_evaluation_step") -def test_manager_enter_collecting_sets_state_and_arms_next_window(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.steps = 11 - manager.step_interval = 20 - counter = torch.ones(4, dtype=torch.int64) - impl = _test_moe_impl( - eplb=True, - route_counter=counter, - recording=True, - num_logical_experts=4, - world_size=1, - ) - manager._eplb_impls = [impl] - - manager._enter_collecting() - - assert torch.equal(counter, torch.ones(4, dtype=torch.int64)) - assert impl.recording - assert manager.state is manager_module.EPLBManagerState.COLLECTING - assert manager.next_evaluation_step == 31 - - def test_manager_step_uses_explicit_state_instead_of_pending_work(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.TRANSFERRING @@ -1526,16 +1483,22 @@ def test_manager_enters_transferring_state_with_planned_work(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) starts = [] + manager.state = manager_module.EPLBManagerState.EVALUATING manager.global_rank = 1 - manager._start_next_transfer = lambda: starts.append(manager.pending_transfer_infos[0]) - - manager._enter_transferring( + manager._evaluation = Future() + manager._evaluation.set_result( { + "kind": "planned", "placement": [[[2], [3]]], "metadata": {0: torch.tensor([1])}, "transfer_infos": [transfer_info], } ) + manager._evaluation_ready_on_all_ranks = lambda: True + manager._publish_expert_load_metric = lambda _result: None + manager._start_next_transfer = lambda: starts.append(manager.pending_transfer_infos[0]) + + manager._step_evaluating() assert manager.state is manager_module.EPLBManagerState.TRANSFERRING assert manager.pending_transfer_infos == [transfer_info] @@ -1549,9 +1512,10 @@ def test_manager_step_waits_for_all_evaluation_results(monkeypatch): manager._evaluation = Future() manager.control_group = object() manager._control_status_buffer = torch.empty(1, dtype=torch.int32) + manager.steps = 11 + manager.step_interval = 20 calls = [] manager._publish_expert_load_metric = lambda _result: calls.append("publish") - manager._enter_collecting = lambda: calls.append("collect") manager.global_rank = 1 monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(0)) @@ -1561,11 +1525,13 @@ def test_manager_step_waits_for_all_evaluation_results(monkeypatch): manager._evaluation.set_result({"kind": "no_improvement"}) monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) manager.step() - assert calls == ["publish", "collect"] + assert calls == ["publish"] + assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert manager.next_evaluation_step == 31 assert not hasattr(manager, "_evaluation") -def test_manager_finish_rebalance_releases_transferring_state(): +def test_manager_complete_rebalance_releases_transferring_state(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) target_placement = torch.tensor([[[2], [3]]]) manager.global_rank = 1 @@ -1578,10 +1544,10 @@ def test_manager_finish_rebalance_releases_transferring_state(): manager.target_metadata = {} manager.rebalance_started_at = time.time() - manager._finish_rebalance() + elapsed = manager._complete_rebalance() assert manager.current_placement is target_placement - assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert elapsed >= 0 for attribute in ( "pending_transfer_infos", "completed_layer_transfers", @@ -1789,6 +1755,14 @@ def test_manager_constructs_pinned_memory_transfer(monkeypatch): monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) monkeypatch.setattr(manager_module, "get_eplb_step_interval", lambda: 20) monkeypatch.setattr(manager_module, "get_eplb_rebalance_gain_threshold", lambda: 0.07) + monkeypatch.setattr(manager_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=1234)) + metric_client = SimpleNamespace() + metric_client_ports = [] + monkeypatch.setattr( + manager_module, + "MetricClient", + lambda port: (metric_client_ports.append(port) or metric_client), + ) def new_group(*args, **kwargs): new_group_calls.append((args, kwargs)) @@ -1820,6 +1794,8 @@ def new_group(*args, **kwargs): assert transfer_calls == [([weight], groups[2], 0, transfer_info)] assert transfer_starts == [True] assert manager.planner.rebalance_gain_threshold == 0.07 + assert manager.metric_client is metric_client + assert metric_client_ports == [1234] assert manager.next_evaluation_step == manager.step_interval assert "planner=GreedyEPLBPlanner" in logs[0] assert weight.fuse_moe_impl.recording From b32ac322ba6d4da3bdf243235fec0759990c815d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 09:07:57 +0000 Subject: [PATCH 33/72] refactor(eplb): simplify manager state transitions --- .../model_infer/mode_backend/eplb_manager.py | 182 ++++++++++----- unit_tests/common/fused_moe/test_eplb.py | 209 ++++++++++++------ 2 files changed, 261 insertions(+), 130 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 75d022f82b..f48a64e52c 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -9,11 +9,11 @@ from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_initial_local_expert_ids, build_logical_to_physical_maps_for_layers, ) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ( EPLBPlanner, + ExpertPlacement, GreedyEPLBPlanner, ) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import ( @@ -38,7 +38,7 @@ logger = init_logger(__name__) EPLB_EXPERT_ALIGNMENT = 128 -EPLB_CONTROL_ERROR = -1 +EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT = 256 EPLB_EXPERT_IMBALANCE_RATIO_METRIC = "lightllm_eplb_topk_expert_imbalance_ratio" @@ -47,6 +47,7 @@ class EPLBManagerState(Enum): COLLECTING = "collecting" EVALUATING = "evaluating" + PLANNING = "planning" TRANSFERRING = "transferring" @@ -55,11 +56,12 @@ class EPLBManager: 状态循环如下: - ``COLLECTING -> EVALUATING -> TRANSFERRING -> COLLECTING`` + ``COLLECTING -> EVALUATING -> PLANNING -> TRANSFERRING -> COLLECTING`` - 当规划器认为无需调整布局时,``EVALUATING`` 会直接回到 + 当收集的平均专家 token 数不足时,``EVALUATING`` 会回到 + ``COLLECTING``;当规划器认为无需调整布局时,``PLANNING`` 会回到 ``COLLECTING``。每次调用 :meth:`step` 最多推进一个状态,耗时的负载 - 评估和权重传输在后台执行,主推理线程只负责轮询和提交结果。 + 聚合、布局规划和权重传输在后台执行,主推理线程只负责轮询和提交结果。 """ def __init__(self, model: TpPartBaseModel) -> None: @@ -82,19 +84,29 @@ def __init__(self, model: TpPartBaseModel) -> None: self.step_interval: int = get_eplb_step_interval() self.steps: int = 0 - # 专家布局与规划器:current_placement 只记录各层的冗余专家槽位。 - initial_local_expert_ids = build_initial_local_expert_ids( - self.num_logical_experts, - self.world_size, - self.num_redundant_experts_per_rank, - ) - initial_redundant_expert_ids = [ - expert_ids[self.num_primary_experts_per_rank :] for expert_ids in initial_local_expert_ids + # 分布式通信:控制面与权重传输使用独立的通信组。 + self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") + self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") + + # 专家布局与规划器:直接以各层 impl 中的实际冗余专家槽位为准。 + # 本 rank 的布局索引为 [layer][redundant_slot]。 + local_redundant_expert_ids_by_layer = [ + impl.local_logics_expert_ids_list[self.num_primary_experts_per_rank :] for impl in self._eplb_impls ] - self.current_placement: torch.Tensor = torch.tensor( - [initial_redundant_expert_ids for _ in weights], - dtype=torch.int64, + + # all_gather 后的布局索引为 [rank][layer][redundant_slot]。 + redundant_expert_ids_by_rank_and_layer: List[List[List[int]]] = [[] for _ in range(self.world_size)] + dist.all_gather_object( + redundant_expert_ids_by_rank_and_layer, + local_redundant_expert_ids_by_layer, + group=self.control_group, ) + + # 转置为规划器使用的 [layer][rank][redundant_slot]。 + self.current_placement: ExpertPlacement = [ + [redundant_expert_ids_by_rank_and_layer[rank][layer_index] for rank in range(self.world_size)] + for layer_index in range(len(weights)) + ] self.planner: EPLBPlanner = GreedyEPLBPlanner( self.world_size, self.num_redundant_experts_per_rank, @@ -102,12 +114,6 @@ def __init__(self, model: TpPartBaseModel) -> None: rebalance_gain_threshold=get_eplb_rebalance_gain_threshold(), ) - # 分布式通信:评估、状态同步和权重传输使用独立的通信组。 - self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") - self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") - self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") - self._control_status_buffer: torch.Tensor = torch.empty(1, dtype=torch.int32) - self.state = EPLBManagerState.COLLECTING self.next_evaluation_step = self.step_interval @@ -129,6 +135,10 @@ def step(self) -> None: self._step_evaluating() return + if self.state is EPLBManagerState.PLANNING: + self._step_planning() + return + if self.state is EPLBManagerState.TRANSFERRING: self._step_transferring() return @@ -138,44 +148,74 @@ def step(self) -> None: # 状态处理:与 step() 的分发顺序保持一致。 def _step_collecting(self) -> None: - """记录一个采样步,并在采样窗口结束后开始评估。""" + """记录一个采样步,并在采样窗口结束后进入评估状态。""" self.steps += 1 if self.steps < self.next_evaluation_step: return - local_load = self._snapshot_local_load() - event = torch.cuda.Event() - event.record(torch.cuda.current_stream()) - evaluation = Future() - del self.next_evaluation_step - self._evaluation = evaluation + self.next_evaluation_step += self.step_interval self.state = EPLBManagerState.EVALUATING - threading.Thread( - target=self._evaluate_after_event, - args=(event, local_load, evaluation), - daemon=True, - ).start() def _step_evaluating(self) -> None: - """等待所有 rank 完成评估,然后进入采样或传输状态。""" - if not self._evaluation_ready_on_all_ranks(): + """启动或等待负载聚合,并根据样本量进入采样或规划状态。""" + if not hasattr(self, "_evaluation"): + local_load = self._snapshot_local_load() + event = torch.cuda.Event() + event.record(torch.cuda.current_stream()) + evaluation = Future() + self._evaluation = evaluation + threading.Thread( + target=self._evaluate_after_event, + args=(event, local_load, evaluation), + daemon=True, + ).start() + return + + if not self._background_work_ready_on_all_ranks(self._evaluation, "evaluation"): return result = self._evaluation.result() del self._evaluation self._publish_expert_load_metric(result) + if result["average_tokens_per_expert"] < EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT: + if self.global_rank == 0: + logger.info( + "eplb continue collecting average_tokens_per_expert=%.2f threshold=%s", + result["average_tokens_per_expert"], + EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT, + ) + self.state = EPLBManagerState.COLLECTING + return + + planning = Future() + self._planning = planning + self.state = EPLBManagerState.PLANNING + threading.Thread( + target=self._plan_after_evaluation, + args=(result["global_load"], planning), + daemon=True, + ).start() + + def _step_planning(self) -> None: + """等待布局规划,并进入采样或传输状态。""" + if not self._background_work_ready_on_all_ranks(self._planning, "planning"): + return + + result = self._planning.result() + del self._planning if result["kind"] != "planned": if self.global_rank == 0: logger.info("eplb skip rearrangement kind=%s", result["kind"]) self.state = EPLBManagerState.COLLECTING - self.next_evaluation_step = self.steps + self.step_interval return pending_transfer_infos = list(result["transfer_infos"]) if not pending_transfer_infos: raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") - self.target_placement = torch.tensor(result["placement"], dtype=torch.int64) + self.target_placement: ExpertPlacement = [ + [list(expert_ids) for expert_ids in layer_placement] for layer_placement in result["placement"] + ] self.target_metadata = result["metadata"] self.pending_transfer_infos = pending_transfer_infos self.completed_layer_transfers: List[PinnedMemoryEPLBTransfer] = [] @@ -214,7 +254,6 @@ def _step_transferring(self) -> None: elapsed = self._complete_rebalance() self.state = EPLBManagerState.COLLECTING - self.next_evaluation_step = self.steps + self.step_interval if self.global_rank == 0: logger.info( "eplb completed wall_time=%.2fs", @@ -223,18 +262,17 @@ def _step_transferring(self) -> None: # 评估阶段内部实现。 - def _evaluation_ready_on_all_ranks(self) -> bool: - ready = self._evaluation.done() - error = self._evaluation.exception() if ready else None - status = EPLB_CONTROL_ERROR if error is not None else int(ready) - global_evaluation_status = self._control_status_buffer.fill_(status) - dist.all_reduce(global_evaluation_status, op=dist.ReduceOp.MIN, group=self.control_group) - global_evaluation_status_value = int(global_evaluation_status.item()) - if global_evaluation_status_value < 0: + def _background_work_ready_on_all_ranks(self, work: Future, phase: str) -> bool: + if not work.done(): + return False + error = work.exception() + failed_by_rank = [False] * self.world_size + dist.all_gather_object(failed_by_rank, error is not None, group=self.control_group) + if any(failed_by_rank): if error is not None: - raise RuntimeError("EPLB evaluation failed on this rank") from error - raise RuntimeError("EPLB evaluation failed on another rank") - return bool(global_evaluation_status_value) + raise RuntimeError(f"EPLB {phase} failed on this rank") from error + raise RuntimeError(f"EPLB {phase} failed on another rank") + return True def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: if self.global_rank != 0: @@ -260,14 +298,25 @@ def _evaluate_after_event( torch.cuda.set_device(local_load.device) event.synchronize() global_load = local_load.cpu() - dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) + dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.control_group) + evaluation.set_result( + { + "global_load": global_load, + "average_tokens_per_expert": _average_tokens_per_expert(global_load), + "expert_imbalance_ratio": _expert_load_imbalance_ratio(global_load), + } + ) + except BaseException as exc: + evaluation.set_exception(exc) + + def _plan_after_evaluation(self, global_load: torch.Tensor, planning: Future) -> None: + try: result = self._plan_and_broadcast(global_load) - result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) if result["kind"] == "planned": result["metadata"], result["transfer_infos"] = self._build_rebalance_data(result) - evaluation.set_result(result) + planning.set_result(result) except BaseException as exc: - evaluation.set_exception(exc) + planning.set_exception(exc) def _plan_and_broadcast(self, global_load: torch.Tensor) -> Dict[str, Any]: result: Optional[Dict[str, Any]] = None @@ -276,13 +325,13 @@ def _plan_and_broadcast(self, global_load: torch.Tensor) -> Dict[str, Any]: try: result = self.planner.plan( global_load.tolist(), - self.current_placement.tolist(), + self.current_placement, ).as_dict() except BaseException as exc: local_error = exc result = {"kind": "error", "message": f"{type(exc).__name__}: {exc}"} values = [result] - dist.broadcast_object_list(values, src=0, group=self.evaluation_group) + dist.broadcast_object_list(values, src=0, group=self.control_group) result = values[0] if result["kind"] == "error": if local_error is not None: @@ -321,7 +370,7 @@ def _build_rebalance_data(self, result: Dict[str, Any]) -> Tuple[Dict[int, torch dtype=torch.int32, ) for changed_layer_offset, layer_index in enumerate(changed_layer_indices): - current_layer_placement = self.current_placement[layer_index].tolist() + current_layer_placement = self.current_placement[layer_index] target_layer_placement = result["placement"][layer_index] metadata_by_layer[layer_index] = logical_to_physical_maps[changed_layer_offset] planned_transfers.extend( @@ -339,9 +388,13 @@ def _build_rebalance_data(self, result: Dict[str, Any]) -> Tuple[Dict[int, torch def _active_transfer_finished_on_all_ranks(self) -> bool: """仅当所有 rank 都完成当前传输时返回 ``True``。""" - all_ranks_finished = self._control_status_buffer.fill_(int(self.active_transfer.is_finished())) - dist.all_reduce(all_ranks_finished, op=dist.ReduceOp.MIN, group=self.control_group) - return bool(all_ranks_finished.item()) + finished_by_rank = [False] * self.world_size + dist.all_gather_object( + finished_by_rank, + self.active_transfer.is_finished(), + group=self.control_group, + ) + return all(finished_by_rank) def _complete_active_transfer(self) -> EPLBTransferInfo: """将当前传输从待处理队列移动到本层的已完成列表。""" @@ -390,7 +443,7 @@ def _complete_rebalance(self) -> float: def _commit_transferred_layer(self, layer_index: int) -> None: """在主推理线程中同步发布一层权重和路由 metadata。""" - target_redundant_expert_ids: List[int] = self.target_placement[layer_index, self.global_rank].tolist() + target_redundant_expert_ids: List[int] = self.target_placement[layer_index][self.global_rank] for transfer in self.completed_layer_transfers: transfer_info: EPLBTransferInfo = transfer.transfer_info if transfer_info.dest_rank != self.global_rank: @@ -419,6 +472,13 @@ def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: return float(ratios.mean().item()) +def _average_tokens_per_expert(global_load: torch.Tensor) -> float: + """Average token count for each logical expert in each layer.""" + if global_load.ndim != 2: + raise ValueError("global_load must be [layers, logical_experts]") + return float(global_load.to(torch.float64).mean().item()) + + def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: weights_by_id: Dict[int, FusedMoeWeight] = {} for layer in model.trans_layers_weight: diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 69b46d4227..ccf5b4ca8e 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -71,6 +71,7 @@ def _test_moe_impl( num_primary_experts_per_rank=num_logical_experts // world_size, num_total_physical_experts=(num_logical_experts + world_size * num_redundant_experts_per_rank), num_redundant_experts_per_rank=num_redundant_experts_per_rank, + local_logics_expert_ids_list=list(range(num_logical_experts // world_size + num_redundant_experts_per_rank)), logical_to_physical_map=logical_to_physical_map, route_counter=route_counter, recording=recording, @@ -589,8 +590,8 @@ def test_manager_delegates_distribution_planning_to_planner_class(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.global_rank = 0 manager.world_size = 2 - manager.evaluation_group = object() - manager.current_placement = torch.tensor([[[1]]]) + manager.control_group = object() + manager.current_placement = [[[1]]] logical_load = torch.tensor([[10, 20]]) calls = [] @@ -611,8 +612,8 @@ def as_dict(self): assert result == {"kind": "no_improvement"} assert len(calls) == 1 assert calls[0][0] == logical_load.tolist() - assert calls[0][1] == manager.current_placement.tolist() - assert broadcasts == [([result], {"src": 0, "group": manager.evaluation_group})] + assert calls[0][1] == manager.current_placement + assert broadcasts == [([result], {"src": 0, "group": manager.control_group})] def test_expert_load_imbalance_ratio_averages_layer_ratios(): @@ -722,8 +723,8 @@ def test_manager_evaluation_collective_sums_rank_loads(monkeypatch): manager.step_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 - manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0) - manager.evaluation_group = object() + manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).tolist() + manager.control_group = object() local = torch.full((1, 4), 100, dtype=torch.int64) seen = {} @@ -736,12 +737,6 @@ def all_reduce(tensor, **kwargs): monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) - def plan_and_broadcast(global_load): - seen["global_load"] = global_load.clone() - return {"kind": "no_improvement"} - - manager._plan_and_broadcast = plan_and_broadcast - evaluation = Future() manager._evaluate_after_event( type("Event", (), {"synchronize": lambda self: None})(), @@ -750,14 +745,15 @@ def plan_and_broadcast(global_load): ) result = evaluation.result() - assert seen["group"] is manager.evaluation_group + assert seen["group"] is manager.control_group assert seen["before"].shape == (1, 4) assert torch.equal(seen["before"], local) - assert torch.equal(seen["global_load"], torch.full_like(local, 200)) + assert torch.equal(result["global_load"], torch.full_like(local, 200)) + assert result["average_tokens_per_expert"] == 200.0 assert result["expert_imbalance_ratio"] == 1.0 -def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_call( +def test_manager_planning_builds_improved_metadata_in_one_multilayer_call( monkeypatch, ): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) @@ -785,8 +781,8 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c manager.step_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 - manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() - manager.evaluation_group = object() + manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone().tolist() + manager.control_group = object() local = torch.full((3, 4), 100, dtype=torch.int64) planned_placement = torch.tensor( [ @@ -801,8 +797,6 @@ def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_c "placement": planned_placement.tolist(), "changed_layers": [True, False, True], } - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) - monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) calls = [] original_build_maps_for_layers = manager_module.build_logical_to_physical_maps_for_layers @@ -817,13 +811,9 @@ def build_maps_for_layers(*args, **kwargs): build_maps_for_layers, ) - evaluation = Future() - manager._evaluate_after_event( - type("Event", (), {"synchronize": lambda self: None})(), - local, - evaluation, - ) - result = evaluation.result() + planning = Future() + manager._plan_after_evaluation(local, planning) + result = planning.result() assert calls == [(2, 4, 2)] metadata = result["metadata"] @@ -1245,7 +1235,7 @@ def test_manager_commits_completed_transfer_rows_by_planned_destination_index(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.global_rank = 0 manager.num_primary_experts_per_rank = 3 - manager.target_placement = torch.tensor([[[7, 6]]]) + manager.target_placement = [[[7, 6]]] manager._eplb_impls = [SimpleNamespace(local_logics_expert_ids_list=local_expert_ids)] manager._commit_layer_metadata = lambda _layer: None manager.completed_layer_transfers = [ @@ -1282,7 +1272,7 @@ def is_finished(self): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.active_transfer = Transfer(info0) manager.control_group = object() - manager._control_status_buffer = torch.empty(1, dtype=torch.int32) + manager.world_size = 2 manager.pending_transfer_infos = [info0, info0b, info1] manager.completed_layer_transfers = [] manager.global_rank = 1 @@ -1306,14 +1296,17 @@ def wait_stream(self, stream): monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: CurrentStream()) monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) - def set_global_ready(count): - return lambda tensor, **_kwargs: tensor.fill_(count) + def set_global_ready(ready): + def all_gather_object(output, _local_ready, **_kwargs): + output[:] = [ready] * manager.world_size + + return all_gather_object - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(0)) + monkeypatch.setattr(manager_module.dist, "all_gather_object", set_global_ready(False)) manager._step_transferring() assert committed == [] - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) + monkeypatch.setattr(manager_module.dist, "all_gather_object", set_global_ready(True)) manager._step_transferring() assert committed == [] assert operations == [] @@ -1334,36 +1327,40 @@ def set_global_ready(count): manager._step_transferring() -def test_evaluation_ready_gate_propagates_local_and_remote_errors(monkeypatch): +def test_background_work_ready_reports_pending_completion_and_errors(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.control_group = object() - manager._control_status_buffer = torch.empty(1, dtype=torch.int32) + manager.world_size = 2 + manager._evaluation = Future() + assert not manager._background_work_ready_on_all_ranks(manager._evaluation, "evaluation") + manager._evaluation.set_exception(RuntimeError("evaluation boom")) statuses = [] - def retain_local_error(tensor, **_kwargs): - statuses.append(int(tensor.item())) + def retain_local_error(output, local_failed, **_kwargs): + statuses.append(local_failed) + output[:] = [local_failed, False] - monkeypatch.setattr(manager_module.dist, "all_reduce", retain_local_error) + monkeypatch.setattr(manager_module.dist, "all_gather_object", retain_local_error) with pytest.raises(RuntimeError, match="EPLB evaluation failed on this rank") as exc_info: - manager._evaluation_ready_on_all_ranks() + manager._background_work_ready_on_all_ranks(manager._evaluation, "evaluation") assert isinstance(exc_info.value.__cause__, RuntimeError) assert str(exc_info.value.__cause__) == "evaluation boom" - assert statuses == [manager_module.EPLB_CONTROL_ERROR] + assert statuses == [True] manager._evaluation = Future() manager._evaluation.set_result({"kind": "no_improvement"}) statuses.clear() - def remote_error(tensor, **_kwargs): - statuses.append(int(tensor.item())) - tensor.fill_(manager_module.EPLB_CONTROL_ERROR) + def remote_error(output, local_failed, **_kwargs): + statuses.append(local_failed) + output[:] = [local_failed, True] - monkeypatch.setattr(manager_module.dist, "all_reduce", remote_error) + monkeypatch.setattr(manager_module.dist, "all_gather_object", remote_error) with pytest.raises(RuntimeError, match="EPLB evaluation failed on another rank"): - manager._evaluation_ready_on_all_ranks() - assert statuses == [1] + manager._background_work_ready_on_all_ranks(manager._evaluation, "evaluation") + assert statuses == [False] @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @@ -1392,18 +1389,22 @@ def is_finished(self): transfer_info = EPLBTransferInfo(0, 0, 2, 0, 0) manager.active_transfer = Transfer(live, received, transfer_info) manager.control_group = object() - manager._control_status_buffer = torch.empty(1, dtype=torch.int32) + manager.world_size = 1 manager.pending_transfer_infos = [transfer_info] manager.completed_layer_transfers = [] manager.num_primary_experts_per_rank = 0 manager.global_rank = 0 - manager.target_placement = torch.tensor([[[2]]]) + manager.target_placement = [[[2]]] manager._eplb_impls = [SimpleNamespace(local_logics_expert_ids_list=[1])] manager._commit_layer_metadata = lambda _layer: None manager._complete_rebalance = lambda: 0.0 manager.steps = 0 manager.step_interval = 20 - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) + monkeypatch.setattr( + manager_module.dist, + "all_gather_object", + lambda output, local_ready, **_kwargs: output.__setitem__(slice(None), [local_ready]), + ) try: g_infer_context.overlap_stream = source_stream @@ -1435,7 +1436,7 @@ def test_manager_step_advances_inflight_transfer(): assert calls == ["transfer"] -def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): +def test_manager_starts_evaluation_only_after_entering_evaluating_state(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) local_load = torch.tensor([[1, 2]], dtype=torch.int64) event = SimpleNamespace(record=lambda stream: None) @@ -1460,9 +1461,17 @@ def test_manager_uses_one_fixed_full_sampling_window(monkeypatch): assert manager.state is manager_module.EPLBManagerState.COLLECTING manager.step() + assert manager.state is manager_module.EPLBManagerState.EVALUATING + assert thread_args == [] + assert not hasattr(manager, "_evaluation") + assert manager.next_evaluation_step == 6 + + manager.step() + assert manager.state is manager_module.EPLBManagerState.EVALUATING assert thread_args[0][1] is local_load - assert not hasattr(manager, "next_evaluation_step") + assert thread_args[0][2] is manager._evaluation + assert manager.next_evaluation_step == 6 def test_manager_step_uses_explicit_state_instead_of_pending_work(): @@ -1483,10 +1492,10 @@ def test_manager_enters_transferring_state_with_planned_work(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) starts = [] - manager.state = manager_module.EPLBManagerState.EVALUATING + manager.state = manager_module.EPLBManagerState.PLANNING manager.global_rank = 1 - manager._evaluation = Future() - manager._evaluation.set_result( + manager._planning = Future() + manager._planning.set_result( { "kind": "planned", "placement": [[[2], [3]]], @@ -1494,11 +1503,10 @@ def test_manager_enters_transferring_state_with_planned_work(): "transfer_infos": [transfer_info], } ) - manager._evaluation_ready_on_all_ranks = lambda: True - manager._publish_expert_load_metric = lambda _result: None + manager._background_work_ready_on_all_ranks = lambda _work, _phase: True manager._start_next_transfer = lambda: starts.append(manager.pending_transfer_infos[0]) - manager._step_evaluating() + manager._step_planning() assert manager.state is manager_module.EPLBManagerState.TRANSFERRING assert manager.pending_transfer_infos == [transfer_info] @@ -1511,19 +1519,29 @@ def test_manager_step_waits_for_all_evaluation_results(monkeypatch): manager.state = manager_module.EPLBManagerState.EVALUATING manager._evaluation = Future() manager.control_group = object() - manager._control_status_buffer = torch.empty(1, dtype=torch.int32) + manager.world_size = 2 manager.steps = 11 manager.step_interval = 20 + manager.next_evaluation_step = 31 calls = [] manager._publish_expert_load_metric = lambda _result: calls.append("publish") manager.global_rank = 1 - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(0)) manager.step() assert calls == [] - manager._evaluation.set_result({"kind": "no_improvement"}) - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) + manager._evaluation.set_result( + { + "global_load": torch.full((1, 4), 255, dtype=torch.int64), + "average_tokens_per_expert": 255.0, + "expert_imbalance_ratio": 1.0, + } + ) + monkeypatch.setattr( + manager_module.dist, + "all_gather_object", + lambda output, local_failed, **_kwargs: output.__setitem__(slice(None), [local_failed, local_failed]), + ) manager.step() assert calls == ["publish"] assert manager.state is manager_module.EPLBManagerState.COLLECTING @@ -1531,9 +1549,57 @@ def test_manager_step_waits_for_all_evaluation_results(monkeypatch): assert not hasattr(manager, "_evaluation") +def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + global_load = torch.full((1, 4), 256, dtype=torch.int64) + worker = SimpleNamespace(start=lambda: None) + thread_args = [] + manager.state = manager_module.EPLBManagerState.EVALUATING + manager.global_rank = 1 + manager._evaluation = Future() + manager._evaluation.set_result( + { + "global_load": global_load, + "average_tokens_per_expert": 256.0, + "expert_imbalance_ratio": 1.0, + } + ) + manager._background_work_ready_on_all_ranks = lambda _work, _phase: True + manager._publish_expert_load_metric = lambda _result: None + monkeypatch.setattr( + manager_module.threading, + "Thread", + lambda **kwargs: (thread_args.append(kwargs["args"]) or worker), + ) + + manager.step() + + assert manager.state is manager_module.EPLBManagerState.PLANNING + assert thread_args[0][0] is global_load + assert thread_args[0][1] is manager._planning + assert not hasattr(manager, "_evaluation") + + +def test_manager_planning_without_changes_returns_to_collecting(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.state = manager_module.EPLBManagerState.PLANNING + manager.global_rank = 1 + manager.steps = 11 + manager.step_interval = 20 + manager.next_evaluation_step = 31 + manager._planning = Future() + manager._planning.set_result({"kind": "no_improvement"}) + manager._background_work_ready_on_all_ranks = lambda _work, _phase: True + manager.step() + + assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert manager.next_evaluation_step == 31 + assert not hasattr(manager, "_planning") + + def test_manager_complete_rebalance_releases_transferring_state(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - target_placement = torch.tensor([[[2], [3]]]) + target_placement = [[[2], [3]]] manager.global_rank = 1 manager.steps = 10 manager.step_interval = 20 @@ -1748,7 +1814,7 @@ def test_manager_constructs_pinned_memory_transfer(monkeypatch): )() transfer_starts = [] transfer = SimpleNamespace(start=lambda: transfer_starts.append(True)) - groups = [object(), object(), object()] + groups = [object(), object()] new_group_calls = [] monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) @@ -1769,6 +1835,13 @@ def new_group(*args, **kwargs): return groups[len(new_group_calls) - 1] monkeypatch.setattr(manager_module.dist, "new_group", new_group) + all_gather_calls = [] + + def all_gather_object(output, local_redundant_expert_ids_by_layer, group): + all_gather_calls.append((local_redundant_expert_ids_by_layer, group)) + output[:] = [local_redundant_expert_ids_by_layer, [[0, 1]]] + + monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) transfer_calls = [] monkeypatch.setattr( manager_module, @@ -1782,18 +1855,16 @@ def new_group(*args, **kwargs): assert not hasattr(manager, "active_transfer") assert not hasattr(manager, "target_placement") assert manager.state is manager_module.EPLBManagerState.COLLECTING - assert ( - manager.evaluation_group, - manager.control_group, - manager.transfer_group, - ) == tuple(groups) - assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 + assert (manager.control_group, manager.transfer_group) == tuple(groups) + assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 2 + assert all_gather_calls == [([[2, 3]], groups[0])] transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) manager.pending_transfer_infos = [transfer_info] manager._start_next_transfer() - assert transfer_calls == [([weight], groups[2], 0, transfer_info)] + assert transfer_calls == [([weight], groups[1], 0, transfer_info)] assert transfer_starts == [True] assert manager.planner.rebalance_gain_threshold == 0.07 + assert manager.current_placement == [[[2, 3], [0, 1]]] assert manager.metric_client is metric_client assert metric_client_ports == [1234] assert manager.next_evaluation_step == manager.step_interval From 1990124dc1ee307823d2dc960bf4872569c4ec53 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 09:42:19 +0000 Subject: [PATCH 34/72] refactor(eplb): separate evaluation and planning phases --- .../model_infer/mode_backend/eplb_manager.py | 103 +++------ unit_tests/common/fused_moe/test_eplb.py | 205 +++++++++--------- 2 files changed, 133 insertions(+), 175 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index f48a64e52c..5f95897c70 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -60,8 +60,8 @@ class EPLBManager: 当收集的平均专家 token 数不足时,``EVALUATING`` 会回到 ``COLLECTING``;当规划器认为无需调整布局时,``PLANNING`` 会回到 - ``COLLECTING``。每次调用 :meth:`step` 最多推进一个状态,耗时的负载 - 聚合、布局规划和权重传输在后台执行,主推理线程只负责轮询和提交结果。 + ``COLLECTING``。每次调用 :meth:`step` 最多推进一个状态,布局规划和 + 权重传输在后台执行,主推理线程负责评估、轮询和提交结果。 """ def __init__(self, model: TpPartBaseModel) -> None: @@ -157,52 +157,55 @@ def _step_collecting(self) -> None: self.state = EPLBManagerState.EVALUATING def _step_evaluating(self) -> None: - """启动或等待负载聚合,并根据样本量进入采样或规划状态。""" - if not hasattr(self, "_evaluation"): - local_load = self._snapshot_local_load() - event = torch.cuda.Event() - event.record(torch.cuda.current_stream()) - evaluation = Future() - self._evaluation = evaluation - threading.Thread( - target=self._evaluate_after_event, - args=(event, local_load, evaluation), - daemon=True, - ).start() - return + """将负载复制到 CPU,并根据全局样本量进入采样或规划状态。""" + counters = [impl.route_counter for impl in self._eplb_impls] + if any(counter.ndim != 1 or counter.shape[0] != self.num_logical_experts for counter in counters): + raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") - if not self._background_work_ready_on_all_ranks(self._evaluation, "evaluation"): - return + # 将各层累计的路由计数复制到 CPU,后续规划统一使用这份快照。 + local_load = torch.stack([counter.detach().cpu() for counter in counters]) - result = self._evaluation.result() - del self._evaluation - self._publish_expert_load_metric(result) - if result["average_tokens_per_expert"] < EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT: + # 汇集各 rank 的 token 总数,判断当前统计量是否足以进行布局规划。 + token_count_by_rank = [0] * self.world_size + dist.all_gather_object( + token_count_by_rank, + int(local_load.sum().item()), + group=self.control_group, + ) + average_tokens_per_expert = sum(token_count_by_rank) / local_load.numel() + if average_tokens_per_expert < EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT: if self.global_rank == 0: logger.info( "eplb continue collecting average_tokens_per_expert=%.2f threshold=%s", - result["average_tokens_per_expert"], + average_tokens_per_expert, EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT, ) self.state = EPLBManagerState.COLLECTING return - planning = Future() - self._planning = planning + self._local_load = local_load self.state = EPLBManagerState.PLANNING - threading.Thread( - target=self._plan_after_evaluation, - args=(result["global_load"], planning), - daemon=True, - ).start() def _step_planning(self) -> None: - """等待布局规划,并进入采样或传输状态。""" + """启动或等待全局负载规划,并进入采样或传输状态。""" + if not hasattr(self, "_planning"): + local_load = self._local_load + del self._local_load + planning = Future() + self._planning = planning + threading.Thread( + target=self._plan, + args=(local_load, planning), + daemon=True, + ).start() + return + if not self._background_work_ready_on_all_ranks(self._planning, "planning"): return result = self._planning.result() del self._planning + self._publish_expert_load_metric(result) if result["kind"] != "planned": if self.global_rank == 0: logger.info("eplb skip rearrangement kind=%s", result["kind"]) @@ -279,39 +282,12 @@ def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: return self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) - def _snapshot_local_load(self) -> torch.Tensor: - """在当前计算流中复制计数快照,并立即开始下一个统计窗口。""" - counters = [impl.route_counter for impl in self._eplb_impls] - if any(counter.ndim != 1 or counter.shape[0] != self.num_logical_experts for counter in counters): - raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") - local_load = torch.stack(counters) - torch._foreach_zero_(counters) - return local_load - - def _evaluate_after_event( - self, - event: torch.cuda.Event, - local_load: torch.Tensor, - evaluation: Future, - ) -> None: + def _plan(self, local_load: torch.Tensor, planning: Future) -> None: try: - torch.cuda.set_device(local_load.device) - event.synchronize() - global_load = local_load.cpu() + global_load = local_load.clone() dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.control_group) - evaluation.set_result( - { - "global_load": global_load, - "average_tokens_per_expert": _average_tokens_per_expert(global_load), - "expert_imbalance_ratio": _expert_load_imbalance_ratio(global_load), - } - ) - except BaseException as exc: - evaluation.set_exception(exc) - - def _plan_after_evaluation(self, global_load: torch.Tensor, planning: Future) -> None: - try: result = self._plan_and_broadcast(global_load) + result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) if result["kind"] == "planned": result["metadata"], result["transfer_infos"] = self._build_rebalance_data(result) planning.set_result(result) @@ -472,13 +448,6 @@ def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: return float(ratios.mean().item()) -def _average_tokens_per_expert(global_load: torch.Tensor) -> float: - """Average token count for each logical expert in each layer.""" - if global_load.ndim != 2: - raise ValueError("global_load must be [layers, logical_experts]") - return float(global_load.to(torch.float64).mean().item()) - - def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: weights_by_id: Dict[int, FusedMoeWeight] = {} for layer in model.trans_layers_weight: diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index ccf5b4ca8e..a65cbbeb5a 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -549,41 +549,40 @@ def test_transfer_plan_respects_explicit_target_slots(): } -def test_manager_snapshots_and_resets_logical_route_counters(): +def test_manager_evaluating_copies_route_counters_to_cpu_without_modifying_them(monkeypatch): counters = [ torch.tensor([10, 11], dtype=torch.int64), torch.tensor([40, 41], dtype=torch.int64), ] manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "fuse_moe_impl": _test_moe_impl( - eplb=True, - route_counter=counter, - recording=False, - num_logical_experts=2, - world_size=1, - ), - }, - )() + manager.state = manager_module.EPLBManagerState.EVALUATING + manager._eplb_impls = [ + _test_moe_impl( + eplb=True, + route_counter=counter, + recording=False, + num_logical_experts=2, + world_size=1, + ) for counter in counters ] - manager._eplb_impls = [weight.fuse_moe_impl for weight in manager.weights] - manager._weights = manager.weights - for layer_num, weight in enumerate(manager._weights): - weight.layer_num_ = layer_num manager.num_logical_experts = 2 + manager.global_rank = 1 + manager.world_size = 1 + manager.control_group = object() + local_token_counts = [] - samples = manager._snapshot_local_load() + def all_gather_object(output, local_token_count, **_kwargs): + local_token_counts.append(local_token_count) + output[:] = [0] - assert torch.equal( - samples, - torch.tensor([[10, 11], [40, 41]], dtype=torch.int64), - ) - assert all(torch.count_nonzero(counter) == 0 for counter in counters) + monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) + + manager._step_evaluating() + + assert local_token_counts == [102] + assert torch.equal(counters[0], torch.tensor([10, 11], dtype=torch.int64)) + assert torch.equal(counters[1], torch.tensor([40, 41], dtype=torch.int64)) def test_manager_delegates_distribution_planning_to_planner_class(monkeypatch): @@ -701,7 +700,7 @@ def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): assert not hasattr(impl, "recording") -def test_manager_evaluation_collective_sums_rank_loads(monkeypatch): +def test_manager_evaluation_gathers_token_counts_from_all_ranks(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.weights = [ type( @@ -725,32 +724,24 @@ def test_manager_evaluation_collective_sums_rank_loads(monkeypatch): manager.num_redundant_experts_per_rank = 1 manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).tolist() manager.control_group = object() - local = torch.full((1, 4), 100, dtype=torch.int64) + local = torch.full((4,), 100, dtype=torch.int64) + manager._eplb_impls[0].route_counter = local + manager.state = manager_module.EPLBManagerState.EVALUATING seen = {} - def all_reduce(tensor, **kwargs): - seen["before"] = tensor.clone() + def all_gather_object(output, local_token_count, **kwargs): + seen["local_token_count"] = local_token_count seen["group"] = kwargs["group"] # Simulate one other rank contributing the same logical-expert load. - tensor.add_(100) + output[:] = [local_token_count, local_token_count, 0, 0] - monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) - monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) + monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) - evaluation = Future() - manager._evaluate_after_event( - type("Event", (), {"synchronize": lambda self: None})(), - local.clone(), - evaluation, - ) - result = evaluation.result() + manager._step_evaluating() assert seen["group"] is manager.control_group - assert seen["before"].shape == (1, 4) - assert torch.equal(seen["before"], local) - assert torch.equal(result["global_load"], torch.full_like(local, 200)) - assert result["average_tokens_per_expert"] == 200.0 - assert result["expert_imbalance_ratio"] == 1.0 + assert seen["local_token_count"] == 400 + assert manager.state is manager_module.EPLBManagerState.COLLECTING def test_manager_planning_builds_improved_metadata_in_one_multilayer_call( @@ -810,12 +801,14 @@ def build_maps_for_layers(*args, **kwargs): "build_logical_to_physical_maps_for_layers", build_maps_for_layers, ) + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) planning = Future() - manager._plan_after_evaluation(local, planning) + manager._plan(local, planning) result = planning.result() assert calls == [(2, 4, 2)] + assert result["expert_imbalance_ratio"] == 1.0 metadata = result["metadata"] assert 1 not in metadata assert {info.layer_index for info in result["transfer_infos"]} == {0, 2} @@ -1332,10 +1325,10 @@ def test_background_work_ready_reports_pending_completion_and_errors(monkeypatch manager.control_group = object() manager.world_size = 2 - manager._evaluation = Future() - assert not manager._background_work_ready_on_all_ranks(manager._evaluation, "evaluation") + work = Future() + assert not manager._background_work_ready_on_all_ranks(work, "planning") - manager._evaluation.set_exception(RuntimeError("evaluation boom")) + work.set_exception(RuntimeError("planning boom")) statuses = [] def retain_local_error(output, local_failed, **_kwargs): @@ -1343,14 +1336,14 @@ def retain_local_error(output, local_failed, **_kwargs): output[:] = [local_failed, False] monkeypatch.setattr(manager_module.dist, "all_gather_object", retain_local_error) - with pytest.raises(RuntimeError, match="EPLB evaluation failed on this rank") as exc_info: - manager._background_work_ready_on_all_ranks(manager._evaluation, "evaluation") + with pytest.raises(RuntimeError, match="EPLB planning failed on this rank") as exc_info: + manager._background_work_ready_on_all_ranks(work, "planning") assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "evaluation boom" + assert str(exc_info.value.__cause__) == "planning boom" assert statuses == [True] - manager._evaluation = Future() - manager._evaluation.set_result({"kind": "no_improvement"}) + work = Future() + work.set_result({"kind": "no_improvement"}) statuses.clear() def remote_error(output, local_failed, **_kwargs): @@ -1358,8 +1351,8 @@ def remote_error(output, local_failed, **_kwargs): output[:] = [local_failed, True] monkeypatch.setattr(manager_module.dist, "all_gather_object", remote_error) - with pytest.raises(RuntimeError, match="EPLB evaluation failed on another rank"): - manager._background_work_ready_on_all_ranks(manager._evaluation, "evaluation") + with pytest.raises(RuntimeError, match="EPLB planning failed on another rank"): + manager._background_work_ready_on_all_ranks(work, "planning") assert statuses == [False] @@ -1436,25 +1429,25 @@ def test_manager_step_advances_inflight_transfer(): assert calls == ["transfer"] -def test_manager_starts_evaluation_only_after_entering_evaluating_state(monkeypatch): +def test_manager_evaluates_only_after_entering_evaluating_state(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - local_load = torch.tensor([[1, 2]], dtype=torch.int64) - event = SimpleNamespace(record=lambda stream: None) - worker = SimpleNamespace(start=lambda: None) - thread_args = [] + route_counter = torch.tensor([1, 2], dtype=torch.int64) + local_token_counts = [] manager.state = manager_module.EPLBManagerState.COLLECTING + manager.global_rank = 1 manager.steps = 0 manager.step_interval = 3 manager.next_evaluation_step = 3 - manager._snapshot_local_load = lambda: local_load + manager.num_logical_experts = 2 + manager._eplb_impls = [SimpleNamespace(route_counter=route_counter)] + manager.world_size = 1 + manager.control_group = object() - monkeypatch.setattr(manager_module.torch.cuda, "Event", lambda: event) - monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: object()) - monkeypatch.setattr( - manager_module.threading, - "Thread", - lambda **kwargs: (thread_args.append(kwargs["args"]) or worker), - ) + def all_gather_object(output, local_token_count, **_kwargs): + local_token_counts.append(local_token_count) + output[:] = [0] + + monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) manager.step() manager.step() @@ -1462,23 +1455,22 @@ def test_manager_starts_evaluation_only_after_entering_evaluating_state(monkeypa manager.step() assert manager.state is manager_module.EPLBManagerState.EVALUATING - assert thread_args == [] - assert not hasattr(manager, "_evaluation") + assert local_token_counts == [] assert manager.next_evaluation_step == 6 manager.step() - assert manager.state is manager_module.EPLBManagerState.EVALUATING - assert thread_args[0][1] is local_load - assert thread_args[0][2] is manager._evaluation + assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert local_token_counts == [3] assert manager.next_evaluation_step == 6 + assert torch.equal(route_counter, torch.tensor([1, 2], dtype=torch.int64)) def test_manager_step_uses_explicit_state_instead_of_pending_work(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.TRANSFERRING manager.pending_transfer_infos = [EPLBTransferInfo(0, 0, 2, 1, 2)] - manager._evaluation = Future() + manager._planning = Future() calls = [] manager._step_transferring = lambda: calls.append("transfer") manager._step_evaluating = lambda: calls.append("evaluation") @@ -1501,9 +1493,11 @@ def test_manager_enters_transferring_state_with_planned_work(): "placement": [[[2], [3]]], "metadata": {0: torch.tensor([1])}, "transfer_infos": [transfer_info], + "expert_imbalance_ratio": 1.0, } ) manager._background_work_ready_on_all_ranks = lambda _work, _phase: True + manager._publish_expert_load_metric = lambda _result: None manager._start_next_transfer = lambda: starts.append(manager.pending_transfer_infos[0]) manager._step_planning() @@ -1514,58 +1508,45 @@ def test_manager_enters_transferring_state_with_planned_work(): assert starts == [transfer_info] -def test_manager_step_waits_for_all_evaluation_results(monkeypatch): +def test_manager_evaluation_with_insufficient_tokens_returns_to_collecting(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.EVALUATING - manager._evaluation = Future() - manager.control_group = object() - manager.world_size = 2 + manager._eplb_impls = [SimpleNamespace(route_counter=torch.full((4,), 255, dtype=torch.int64))] + manager.num_logical_experts = 4 manager.steps = 11 manager.step_interval = 20 manager.next_evaluation_step = 31 - calls = [] - manager._publish_expert_load_metric = lambda _result: calls.append("publish") manager.global_rank = 1 - - manager.step() - assert calls == [] - - manager._evaluation.set_result( - { - "global_load": torch.full((1, 4), 255, dtype=torch.int64), - "average_tokens_per_expert": 255.0, - "expert_imbalance_ratio": 1.0, - } - ) + manager.world_size = 1 + manager.control_group = object() monkeypatch.setattr( manager_module.dist, "all_gather_object", - lambda output, local_failed, **_kwargs: output.__setitem__(slice(None), [local_failed, local_failed]), + lambda output, local_token_count, **_kwargs: output.__setitem__(slice(None), [local_token_count]), ) + manager.step() - assert calls == ["publish"] + assert manager.state is manager_module.EPLBManagerState.COLLECTING assert manager.next_evaluation_step == 31 - assert not hasattr(manager, "_evaluation") def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - global_load = torch.full((1, 4), 256, dtype=torch.int64) + local_load = torch.full((1, 4), 256, dtype=torch.int64) worker = SimpleNamespace(start=lambda: None) thread_args = [] manager.state = manager_module.EPLBManagerState.EVALUATING manager.global_rank = 1 - manager._evaluation = Future() - manager._evaluation.set_result( - { - "global_load": global_load, - "average_tokens_per_expert": 256.0, - "expert_imbalance_ratio": 1.0, - } + manager.num_logical_experts = 4 + manager._eplb_impls = [SimpleNamespace(route_counter=local_load[0])] + manager.world_size = 1 + manager.control_group = object() + monkeypatch.setattr( + manager_module.dist, + "all_gather_object", + lambda output, local_token_count, **_kwargs: output.__setitem__(slice(None), [local_token_count]), ) - manager._background_work_ready_on_all_ranks = lambda _work, _phase: True - manager._publish_expert_load_metric = lambda _result: None monkeypatch.setattr( manager_module.threading, "Thread", @@ -1575,9 +1556,16 @@ def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): manager.step() assert manager.state is manager_module.EPLBManagerState.PLANNING - assert thread_args[0][0] is global_load + assert torch.equal(manager._local_load, local_load) + assert thread_args == [] + assert not hasattr(manager, "_planning") + + manager.step() + + assert manager.state is manager_module.EPLBManagerState.PLANNING + assert torch.equal(thread_args[0][0], local_load) assert thread_args[0][1] is manager._planning - assert not hasattr(manager, "_evaluation") + assert not hasattr(manager, "_local_load") def test_manager_planning_without_changes_returns_to_collecting(): @@ -1588,8 +1576,9 @@ def test_manager_planning_without_changes_returns_to_collecting(): manager.step_interval = 20 manager.next_evaluation_step = 31 manager._planning = Future() - manager._planning.set_result({"kind": "no_improvement"}) + manager._planning.set_result({"kind": "no_improvement", "expert_imbalance_ratio": 1.0}) manager._background_work_ready_on_all_ranks = lambda _work, _phase: True + manager._publish_expert_load_metric = lambda _result: None manager.step() assert manager.state is manager_module.EPLBManagerState.COLLECTING @@ -1851,7 +1840,7 @@ def all_gather_object(output, local_redundant_expert_ids_by_layer, group): logs = [] monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) manager = manager_module.EPLBManager(type("Model", (), {})()) - assert not hasattr(manager, "_evaluation") + assert not hasattr(manager, "_planning") assert not hasattr(manager, "active_transfer") assert not hasattr(manager, "target_placement") assert manager.state is manager_module.EPLBManagerState.COLLECTING From bc73bcd8db3cd194e7959ae0ae29bf0c4856c853 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 10:11:35 +0000 Subject: [PATCH 35/72] refactor(eplb): isolate asynchronous planning task --- .../model_infer/mode_backend/eplb_manager.py | 129 ++++------- .../model_infer/mode_backend/eplb_plan.py | 81 +++++++ unit_tests/common/fused_moe/test_eplb.py | 219 ++++++++++-------- 3 files changed, 257 insertions(+), 172 deletions(-) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb_plan.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 5f95897c70..b102d00b74 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -1,6 +1,4 @@ -from concurrent.futures import Future from enum import Enum -import threading import time from typing import Any, Dict, List, Optional, Tuple @@ -20,6 +18,7 @@ FusedMoeWeight, ) from lightllm.server.metrics.manager import MetricClient +from lightllm.server.router.model_infer.mode_backend.eplb_plan import EPLBPlanTask from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( EPLBTransferInfo, PinnedMemoryEPLBTransfer, @@ -48,6 +47,7 @@ class EPLBManagerState(Enum): COLLECTING = "collecting" EVALUATING = "evaluating" PLANNING = "planning" + WAIT_PLAN_FINISH = "wait_plan_finish" TRANSFERRING = "transferring" @@ -56,12 +56,12 @@ class EPLBManager: 状态循环如下: - ``COLLECTING -> EVALUATING -> PLANNING -> TRANSFERRING -> COLLECTING`` + ``COLLECTING -> EVALUATING -> PLANNING -> WAIT_PLAN_FINISH -> TRANSFERRING -> COLLECTING`` 当收集的平均专家 token 数不足时,``EVALUATING`` 会回到 - ``COLLECTING``;当规划器认为无需调整布局时,``PLANNING`` 会回到 - ``COLLECTING``。每次调用 :meth:`step` 最多推进一个状态,布局规划和 - 权重传输在后台执行,主推理线程负责评估、轮询和提交结果。 + ``COLLECTING``;当规划器认为无需调整布局时,``WAIT_PLAN_FINISH`` 会 + 回到 ``COLLECTING``。每次调用 :meth:`step` 最多推进一个状态,布局 + 规划和权重传输在后台执行,主推理线程负责评估、轮询和提交结果。 """ def __init__(self, model: TpPartBaseModel) -> None: @@ -139,6 +139,10 @@ def step(self) -> None: self._step_planning() return + if self.state is EPLBManagerState.WAIT_PLAN_FINISH: + self._step_wait_plan_finish() + return + if self.state is EPLBManagerState.TRANSFERRING: self._step_transferring() return @@ -187,24 +191,46 @@ def _step_evaluating(self) -> None: self.state = EPLBManagerState.PLANNING def _step_planning(self) -> None: - """启动或等待全局负载规划,并进入采样或传输状态。""" - if not hasattr(self, "_planning"): - local_load = self._local_load - del self._local_load - planning = Future() - self._planning = planning - threading.Thread( - target=self._plan, - args=(local_load, planning), - daemon=True, - ).start() - return + """汇集全局负载,并由 rank 0 启动异步规划。""" + local_load = self._local_load + del self._local_load + + # 一次分配连续的 [rank][layer][logical_expert] 缓冲区,再沿 rank 维 + # 切出 all_gather 所需的输出 tensor。 + gathered_load = torch.empty( + (self.world_size, *local_load.shape), + dtype=local_load.dtype, + device=local_load.device, + ) + load_by_rank = list(gathered_load.unbind(dim=0)) + dist.all_gather(load_by_rank, local_load, group=self.control_group) + global_load = gathered_load.sum(dim=0) + + self.state = EPLBManagerState.WAIT_PLAN_FINISH + if self.global_rank == 0: + self._plan_task = EPLBPlanTask( + self.planner, + global_load, + self.current_placement, + ) + self._plan_task.start() + + def _step_wait_plan_finish(self) -> None: + """等待 rank 0 完成规划并广播结果。""" + result: Optional[Dict[str, Any]] = None + if self.global_rank == 0 and self._plan_task.is_finished(): + result = self._plan_task.result + assert result is not None - if not self._background_work_ready_on_all_ranks(self._planning, "planning"): + values = [result] + dist.broadcast_object_list(values, src=0, group=self.control_group) + result = values[0] + if result is None: return - result = self._planning.result() - del self._planning + if self.global_rank == 0: + del self._plan_task + self._publish_expert_load_metric(result) if result["kind"] != "planned": if self.global_rank == 0: @@ -212,6 +238,7 @@ def _step_planning(self) -> None: self.state = EPLBManagerState.COLLECTING return + result["metadata"], result["transfer_infos"] = self._build_rebalance_data(result) pending_transfer_infos = list(result["transfer_infos"]) if not pending_transfer_infos: raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") @@ -263,58 +290,11 @@ def _step_transferring(self) -> None: elapsed, ) - # 评估阶段内部实现。 - - def _background_work_ready_on_all_ranks(self, work: Future, phase: str) -> bool: - if not work.done(): - return False - error = work.exception() - failed_by_rank = [False] * self.world_size - dist.all_gather_object(failed_by_rank, error is not None, group=self.control_group) - if any(failed_by_rank): - if error is not None: - raise RuntimeError(f"EPLB {phase} failed on this rank") from error - raise RuntimeError(f"EPLB {phase} failed on another rank") - return True - def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: if self.global_rank != 0: return self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) - def _plan(self, local_load: torch.Tensor, planning: Future) -> None: - try: - global_load = local_load.clone() - dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.control_group) - result = self._plan_and_broadcast(global_load) - result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(global_load) - if result["kind"] == "planned": - result["metadata"], result["transfer_infos"] = self._build_rebalance_data(result) - planning.set_result(result) - except BaseException as exc: - planning.set_exception(exc) - - def _plan_and_broadcast(self, global_load: torch.Tensor) -> Dict[str, Any]: - result: Optional[Dict[str, Any]] = None - local_error: Optional[BaseException] = None - if self.global_rank == 0: - try: - result = self.planner.plan( - global_load.tolist(), - self.current_placement, - ).as_dict() - except BaseException as exc: - local_error = exc - result = {"kind": "error", "message": f"{type(exc).__name__}: {exc}"} - values = [result] - dist.broadcast_object_list(values, src=0, group=self.control_group) - result = values[0] - if result["kind"] == "error": - if local_error is not None: - raise RuntimeError("EPLB planner failed on rank zero") from local_error - raise RuntimeError(f"EPLB planner failed on rank zero: {result['message']}") - return result - def _build_rebalance_data(self, result: Dict[str, Any]) -> Tuple[Dict[int, torch.Tensor], List[EPLBTransferInfo]]: metadata_by_layer: Dict[int, torch.Tensor] = {} planned_transfers: List[EPLBTransferInfo] = [] @@ -435,19 +415,6 @@ def _commit_layer_metadata(self, layer_index: int) -> None: self._eplb_impls[layer_index].logical_to_physical_map.copy_(self.target_metadata[layer_index]) -def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: - """Average each layer's maximum-to-mean logical-expert token ratio.""" - if global_load.ndim != 2: - raise ValueError("global_load must be [layers, logical_experts]") - global_load = global_load.to(torch.float64) - layer_means = global_load.mean(dim=1) - valid_layers = layer_means > 0 - if not torch.any(valid_layers): - return 0.0 - ratios = global_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] - return float(ratios.mean().item()) - - def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: weights_by_id: Dict[int, FusedMoeWeight] = {} for layer in model.trans_layers_weight: diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_plan.py b/lightllm/server/router/model_infer/mode_backend/eplb_plan.py new file mode 100644 index 0000000000..f38b3a3450 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb_plan.py @@ -0,0 +1,81 @@ +"""EPLB 专家布局的异步规划任务。""" + +import os +import threading +from enum import Enum +from typing import Any, Dict, Optional + +import torch + +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ( + EPLBPlanner, + ExpertPlacement, +) +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +class PlanTaskStatus(Enum): + """异步规划任务的生命周期状态。""" + + IDLE = "idle" + RUNNING = "running" + SUCCEEDED = "succeeded" + + +class EPLBPlanTask: + """在后台线程中根据全局专家负载生成新布局。""" + + def __init__( + self, + planner: EPLBPlanner, + global_load: torch.Tensor, + current_placement: ExpertPlacement, + ) -> None: + self.planner = planner + self.global_load = global_load + self.current_placement = current_placement + self.status = PlanTaskStatus.IDLE + self.result: Optional[Dict[str, Any]] = None + self._thread = threading.Thread( + target=self._run, + name="eplb-plan", + daemon=True, + ) + + def start(self) -> None: + """启动异步规划。""" + assert self.status is PlanTaskStatus.IDLE, "EPLB plan task has already been started" + self.status = PlanTaskStatus.RUNNING + self._thread.start() + + def is_finished(self) -> bool: + """返回规划任务是否已经成功完成。""" + return self.status is PlanTaskStatus.SUCCEEDED + + def _run(self) -> None: + try: + result = self.planner.plan( + self.global_load.tolist(), + self.current_placement, + ).as_dict() + result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(self.global_load) + self.result = result + self.status = PlanTaskStatus.SUCCEEDED + except BaseException: + logger.exception("EPLB planning failed") + os._exit(1) + + +def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: + """Average each layer's maximum-to-mean logical-expert token ratio.""" + if global_load.ndim != 2: + raise ValueError("global_load must be [layers, logical_experts]") + global_load = global_load.to(torch.float64) + layer_means = global_load.mean(dim=1) + valid_layers = layer_means > 0 + if not torch.any(valid_layers): + return 0.0 + ratios = global_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] + return float(ratios.mean().item()) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index a65cbbeb5a..a694e01076 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1,6 +1,5 @@ import threading import time -from concurrent.futures import Future from contextlib import nullcontext from types import SimpleNamespace @@ -22,6 +21,9 @@ from lightllm.server.router.model_infer.mode_backend import ( eplb_manager as manager_module, ) +from lightllm.server.router.model_infer.mode_backend import ( + eplb_plan as plan_module, +) from lightllm.server.router.model_infer.mode_backend import ( eplb_transfer as transfer_module, ) @@ -585,12 +587,8 @@ def all_gather_object(output, local_token_count, **_kwargs): assert torch.equal(counters[1], torch.tensor([40, 41], dtype=torch.int64)) -def test_manager_delegates_distribution_planning_to_planner_class(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.global_rank = 0 - manager.world_size = 2 - manager.control_group = object() - manager.current_placement = [[[1]]] +def test_manager_delegates_distribution_planning_to_planner_class(): + current_placement = [[[1]]] logical_load = torch.tensor([[10, 20]]) calls = [] @@ -598,21 +596,36 @@ class Result: def as_dict(self): return {"kind": "no_improvement"} - manager.planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or Result())) - broadcasts = [] - monkeypatch.setattr( - manager_module.dist, - "broadcast_object_list", - lambda values, **kwargs: broadcasts.append((values, kwargs)), - ) + planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or Result())) + task = plan_module.EPLBPlanTask(planner, logical_load, current_placement) - result = manager._plan_and_broadcast(logical_load) + task._run() - assert result == {"kind": "no_improvement"} + assert task.status is plan_module.PlanTaskStatus.SUCCEEDED + assert task.result == {"kind": "no_improvement", "expert_imbalance_ratio": 4 / 3} assert len(calls) == 1 assert calls[0][0] == logical_load.tolist() - assert calls[0][1] == manager.current_placement - assert broadcasts == [([result], {"src": 0, "group": manager.control_group})] + assert calls[0][1] == current_placement + + +def test_plan_task_exits_process_on_failure(monkeypatch): + def fail(_load, _placement): + raise RuntimeError("planning boom") + + task = plan_module.EPLBPlanTask( + SimpleNamespace(plan=fail), + torch.tensor([[10, 20]]), + [[[1]]], + ) + exits = [] + logs = [] + monkeypatch.setattr(plan_module.os, "_exit", exits.append) + monkeypatch.setattr(plan_module.logger, "exception", logs.append) + + task._run() + + assert exits == [1] + assert logs == ["EPLB planning failed"] def test_expert_load_imbalance_ratio_averages_layer_ratios(): @@ -624,7 +637,7 @@ def test_expert_load_imbalance_ratio_averages_layer_ratios(): dtype=torch.int64, ) - ratio = manager_module._expert_load_imbalance_ratio(global_load) + ratio = plan_module._expert_load_imbalance_ratio(global_load) assert ratio == pytest.approx(1.25) @@ -773,8 +786,6 @@ def test_manager_planning_builds_improved_metadata_in_one_multilayer_call( manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone().tolist() - manager.control_group = object() - local = torch.full((3, 4), 100, dtype=torch.int64) planned_placement = torch.tensor( [ [[3], [0], [1], [2]], @@ -783,7 +794,7 @@ def test_manager_planning_builds_improved_metadata_in_one_multilayer_call( ], dtype=torch.int64, ) - manager._plan_and_broadcast = lambda _global_load: { + result = { "kind": "planned", "placement": planned_placement.tolist(), "changed_layers": [True, False, True], @@ -801,17 +812,11 @@ def build_maps_for_layers(*args, **kwargs): "build_logical_to_physical_maps_for_layers", build_maps_for_layers, ) - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) - - planning = Future() - manager._plan(local, planning) - result = planning.result() + metadata, transfer_infos = manager._build_rebalance_data(result) assert calls == [(2, 4, 2)] - assert result["expert_imbalance_ratio"] == 1.0 - metadata = result["metadata"] assert 1 not in metadata - assert {info.layer_index for info in result["transfer_infos"]} == {0, 2} + assert {info.layer_index for info in transfer_infos} == {0, 2} for layer_index in (0, 2): item = metadata[layer_index] expected = build_logical_to_physical_map( @@ -1320,40 +1325,24 @@ def all_gather_object(output, _local_ready, **_kwargs): manager._step_transferring() -def test_background_work_ready_reports_pending_completion_and_errors(monkeypatch): +def test_wait_plan_finish_broadcasts_pending_status(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.state = manager_module.EPLBManagerState.WAIT_PLAN_FINISH + manager.global_rank = 0 manager.control_group = object() - manager.world_size = 2 - - work = Future() - assert not manager._background_work_ready_on_all_ranks(work, "planning") - - work.set_exception(RuntimeError("planning boom")) - statuses = [] - - def retain_local_error(output, local_failed, **_kwargs): - statuses.append(local_failed) - output[:] = [local_failed, False] - - monkeypatch.setattr(manager_module.dist, "all_gather_object", retain_local_error) - with pytest.raises(RuntimeError, match="EPLB planning failed on this rank") as exc_info: - manager._background_work_ready_on_all_ranks(work, "planning") - assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "planning boom" - assert statuses == [True] - - work = Future() - work.set_result({"kind": "no_improvement"}) - statuses.clear() + manager._plan_task = SimpleNamespace( + is_finished=lambda: False, + result=None, + ) + broadcasts = [] - def remote_error(output, local_failed, **_kwargs): - statuses.append(local_failed) - output[:] = [local_failed, True] + def broadcast(values, **_kwargs): + broadcasts.append(values[0]) - monkeypatch.setattr(manager_module.dist, "all_gather_object", remote_error) - with pytest.raises(RuntimeError, match="EPLB planning failed on another rank"): - manager._background_work_ready_on_all_ranks(work, "planning") - assert statuses == [False] + monkeypatch.setattr(manager_module.dist, "broadcast_object_list", broadcast) + manager._step_wait_plan_finish() + assert broadcasts == [None] + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @@ -1470,7 +1459,7 @@ def test_manager_step_uses_explicit_state_instead_of_pending_work(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.TRANSFERRING manager.pending_transfer_infos = [EPLBTransferInfo(0, 0, 2, 1, 2)] - manager._planning = Future() + manager._plan_task = object() calls = [] manager._step_transferring = lambda: calls.append("transfer") manager._step_evaluating = lambda: calls.append("evaluation") @@ -1480,27 +1469,28 @@ def test_manager_step_uses_explicit_state_instead_of_pending_work(): assert calls == ["transfer"] -def test_manager_enters_transferring_state_with_planned_work(): +def test_manager_enters_transferring_state_with_planned_work(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) starts = [] - manager.state = manager_module.EPLBManagerState.PLANNING + manager.state = manager_module.EPLBManagerState.WAIT_PLAN_FINISH manager.global_rank = 1 - manager._planning = Future() - manager._planning.set_result( - { - "kind": "planned", - "placement": [[[2], [3]]], - "metadata": {0: torch.tensor([1])}, - "transfer_infos": [transfer_info], - "expert_imbalance_ratio": 1.0, - } + manager.control_group = object() + result = { + "kind": "planned", + "placement": [[[2], [3]]], + "expert_imbalance_ratio": 1.0, + } + monkeypatch.setattr( + manager_module.dist, + "broadcast_object_list", + lambda values, **_kwargs: values.__setitem__(0, result), ) - manager._background_work_ready_on_all_ranks = lambda _work, _phase: True manager._publish_expert_load_metric = lambda _result: None + manager._build_rebalance_data = lambda _result: ({0: torch.tensor([1])}, [transfer_info]) manager._start_next_transfer = lambda: starts.append(manager.pending_transfer_infos[0]) - manager._step_planning() + manager._step_wait_plan_finish() assert manager.state is manager_module.EPLBManagerState.TRANSFERRING assert manager.pending_transfer_infos == [transfer_info] @@ -1534,10 +1524,9 @@ def test_manager_evaluation_with_insufficient_tokens_returns_to_collecting(monke def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) local_load = torch.full((1, 4), 256, dtype=torch.int64) - worker = SimpleNamespace(start=lambda: None) - thread_args = [] + plan_tasks = [] manager.state = manager_module.EPLBManagerState.EVALUATING - manager.global_rank = 1 + manager.global_rank = 0 manager.num_logical_experts = 4 manager._eplb_impls = [SimpleNamespace(route_counter=local_load[0])] manager.world_size = 1 @@ -1548,42 +1537,90 @@ def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): lambda output, local_token_count, **_kwargs: output.__setitem__(slice(None), [local_token_count]), ) monkeypatch.setattr( - manager_module.threading, - "Thread", - lambda **kwargs: (thread_args.append(kwargs["args"]) or worker), + manager_module.dist, + "all_gather", + lambda output, local, **_kwargs: output[0].copy_(local), ) + class PlanTask: + def __init__(self, planner, global_load, current_placement): + self.planner = planner + self.global_load = global_load + self.current_placement = current_placement + self.started = False + plan_tasks.append(self) + + def start(self): + self.started = True + + manager.planner = object() + manager.current_placement = [[[1]]] + monkeypatch.setattr(manager_module, "EPLBPlanTask", PlanTask) + manager.step() assert manager.state is manager_module.EPLBManagerState.PLANNING assert torch.equal(manager._local_load, local_load) - assert thread_args == [] - assert not hasattr(manager, "_planning") + assert plan_tasks == [] + assert not hasattr(manager, "_plan_task") manager.step() - assert manager.state is manager_module.EPLBManagerState.PLANNING - assert torch.equal(thread_args[0][0], local_load) - assert thread_args[0][1] is manager._planning + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH + assert torch.equal(plan_tasks[0].global_load, local_load) + assert plan_tasks[0].planner is manager.planner + assert plan_tasks[0].current_placement is manager.current_placement + assert plan_tasks[0].started + assert manager._plan_task is plan_tasks[0] assert not hasattr(manager, "_local_load") -def test_manager_planning_without_changes_returns_to_collecting(): +def test_nonzero_rank_waits_without_starting_planner(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.PLANNING manager.global_rank = 1 + manager.world_size = 2 + manager.control_group = object() + manager._local_load = torch.tensor([[1, 2]], dtype=torch.int64) + + def all_gather(output, local, **_kwargs): + output[0].copy_(local) + output[1].copy_(local) + + monkeypatch.setattr(manager_module.dist, "all_gather", all_gather) + + manager.step() + + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH + assert not hasattr(manager, "_plan_task") + assert not hasattr(manager, "_local_load") + + +def test_manager_planning_without_changes_returns_to_collecting(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.state = manager_module.EPLBManagerState.WAIT_PLAN_FINISH + manager.global_rank = 1 + manager.control_group = object() manager.steps = 11 manager.step_interval = 20 manager.next_evaluation_step = 31 - manager._planning = Future() - manager._planning.set_result({"kind": "no_improvement", "expert_imbalance_ratio": 1.0}) - manager._background_work_ready_on_all_ranks = lambda _work, _phase: True manager._publish_expert_load_metric = lambda _result: None + result = None + + def broadcast(values, **_kwargs): + values[0] = result + + monkeypatch.setattr(manager_module.dist, "broadcast_object_list", broadcast) + + manager.step() + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH + + result = {"kind": "no_improvement", "expert_imbalance_ratio": 1.0} manager.step() assert manager.state is manager_module.EPLBManagerState.COLLECTING assert manager.next_evaluation_step == 31 - assert not hasattr(manager, "_planning") + assert not hasattr(manager, "_plan_task") def test_manager_complete_rebalance_releases_transferring_state(): @@ -1840,7 +1877,7 @@ def all_gather_object(output, local_redundant_expert_ids_by_layer, group): logs = [] monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) manager = manager_module.EPLBManager(type("Model", (), {})()) - assert not hasattr(manager, "_planning") + assert not hasattr(manager, "_plan_task") assert not hasattr(manager, "active_transfer") assert not hasattr(manager, "target_placement") assert manager.state is manager_module.EPLBManagerState.COLLECTING From 81bb2fc3a7412c4828aa9ee2843fc76d5ee64539 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 11:51:12 +0000 Subject: [PATCH 36/72] refactor EPLB planning and transfer lifecycle --- .../meta_weights/fused_moe/eplb_planner.py | 65 +--- .../model_infer/mode_backend/eplb_manager.py | 217 +++++------- .../model_infer/mode_backend/eplb_plan.py | 23 +- unit_tests/common/fused_moe/test_eplb.py | 330 ++++++++---------- 4 files changed, 240 insertions(+), 395 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py index bcf43a6b14..4d6c4aeddd 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py @@ -8,9 +8,8 @@ from abc import ABC, abstractmethod from collections import Counter -from dataclasses import dataclass from math import ceil -from typing import Any, Dict, List +from typing import List # [layer][logical expert] @@ -21,35 +20,6 @@ RankLoad = List[List[float]] -@dataclass -class EPLBPlan: - placement: ExpertPlacement - changed_layers: List[bool] - before_rank_load: RankLoad - after_rank_load: RankLoad - reason: str - - @property - def changed(self) -> bool: - return any(self.changed_layers) - - def as_dict(self) -> Dict[str, Any]: - before = _imbalance_summary(self.before_rank_load) - after = _imbalance_summary(self.after_rank_load) - before_critical = sum(max(layer) for layer in self.before_rank_load) - after_critical = sum(max(layer) for layer in self.after_rank_load) - gain = (before_critical - after_critical) / max(before_critical, 1.0) - return { - "kind": self.reason, - "placement": self.placement, - "changed_layers": self.changed_layers, - "before": before, - "after": after, - "rebalance_gain": gain, - "changed_layer_count": sum(self.changed_layers), - } - - class EPLBPlanner(ABC): """Interface for planning redundant expert placement.""" @@ -58,8 +28,8 @@ def plan( self, logical_expert_load: LogicalExpertLoad, current_placement: ExpertPlacement, - ) -> EPLBPlan: - """Plan a concrete ``[layer][rank][redundant slot]`` placement.""" + ) -> ExpertPlacement: + """Return a concrete ``[layer][rank][redundant slot]`` placement.""" class GreedyEPLBPlanner(EPLBPlanner): @@ -90,7 +60,7 @@ def plan( self, logical_expert_load: LogicalExpertLoad, current_placement: ExpertPlacement, - ) -> EPLBPlan: + ) -> ExpertPlacement: """Return a new concrete placement. ``logical_expert_load`` is ``[layer][logical_expert]`` and contains @@ -106,25 +76,14 @@ def plan( candidates = [self._plan_layer(layer_load, current_layer) for layer_load, current_layer in zip(load, current)] candidate_rank_load = self.estimate_rank_load(load, candidates) - changed_layers = [] placement = [] - after_rank_load = [] for layer, candidate in enumerate(candidates): before = max(before_rank_load[layer]) after = max(candidate_rank_load[layer]) gain = (before - after) / max(before, 1.0) changed = candidate != current[layer] and gain > self.rebalance_gain_threshold - changed_layers.append(changed) placement.append(candidate if changed else current[layer]) - after_rank_load.append(candidate_rank_load[layer][:] if changed else before_rank_load[layer][:]) - - return EPLBPlan( - placement=placement, - changed_layers=changed_layers, - before_rank_load=before_rank_load, - after_rank_load=after_rank_load, - reason="planned" if any(changed_layers) else "no_improvement", - ) + return placement def estimate_rank_load( self, @@ -365,17 +324,3 @@ def _validate_inputs( raise ValueError("placement contains an invalid logical expert") self._expert_locations(layer_placement, num_logical_experts) return num_logical_experts - - -def _imbalance_summary(rank_load: RankLoad) -> Dict[str, float]: - values = sorted(value for layer in rank_load for value in layer) - if not values: - return {"max": 0.0, "p95": 0.0, "mean": 0.0, "ratio": 0.0} - mean = sum(values) / len(values) - p95 = values[ceil(0.95 * len(values)) - 1] - return { - "max": max(values), - "p95": p95, - "mean": mean, - "ratio": max(values) / max(mean, 1.0), - } diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index b102d00b74..c36e814d7b 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -1,13 +1,13 @@ from enum import Enum import time -from typing import Any, Dict, List, Optional, Tuple +from typing import Dict, List, Optional import torch import torch.distributed as dist from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_logical_to_physical_maps_for_layers, + build_logical_to_physical_map, ) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ( EPLBPlanner, @@ -205,6 +205,7 @@ def _step_planning(self) -> None: load_by_rank = list(gathered_load.unbind(dim=0)) dist.all_gather(load_by_rank, local_load, group=self.control_group) global_load = gathered_load.sum(dim=0) + self._publish_expert_load_metric(global_load) self.state = EPLBManagerState.WAIT_PLAN_FINISH if self.global_rank == 0: @@ -216,68 +217,66 @@ def _step_planning(self) -> None: self._plan_task.start() def _step_wait_plan_finish(self) -> None: - """等待 rank 0 完成规划并广播结果。""" - result: Optional[Dict[str, Any]] = None + """等待 rank 0 完成规划并广播目标专家排布。""" + placement: Optional[ExpertPlacement] = None if self.global_rank == 0 and self._plan_task.is_finished(): - result = self._plan_task.result - assert result is not None + placement = self._plan_task.result + assert placement is not None - values = [result] + values = [placement] dist.broadcast_object_list(values, src=0, group=self.control_group) - result = values[0] - if result is None: + placement = values[0] + if placement is None: return if self.global_rank == 0: del self._plan_task - self._publish_expert_load_metric(result) - if result["kind"] != "planned": + if placement == self.current_placement: if self.global_rank == 0: - logger.info("eplb skip rearrangement kind=%s", result["kind"]) + logger.info("eplb skip rearrangement because placement is unchanged") self.state = EPLBManagerState.COLLECTING return - result["metadata"], result["transfer_infos"] = self._build_rebalance_data(result) - pending_transfer_infos = list(result["transfer_infos"]) - if not pending_transfer_infos: - raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") - self.target_placement: ExpertPlacement = [ - [list(expert_ids) for expert_ids in layer_placement] for layer_placement in result["placement"] + [list(expert_ids) for expert_ids in layer_placement] for layer_placement in placement + ] + self.pending_transfer_infos = [ + transfer_info + for layer_index, (current_layer, target_layer) in enumerate( + zip(self.current_placement, self.target_placement) + ) + for transfer_info in build_transfer_plan( + current_layer, + target_layer, + layer_index, + self.num_logical_experts, + self.world_size, + ) ] - self.target_metadata = result["metadata"] - self.pending_transfer_infos = pending_transfer_infos - self.completed_layer_transfers: List[PinnedMemoryEPLBTransfer] = [] - self.rebalance_started_at = time.time() + if not self.pending_transfer_infos: + raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") self.state = EPLBManagerState.TRANSFERRING - self._start_next_transfer() if self.global_rank == 0: logger.info( - "eplb started steps=%s max_before=%.4f max_after=%.4f " - "p95_before=%.4f p95_after=%.4f rebalance_gain=%.4f " - "changed_layer_count=%s changed_slot_count=%s", + "eplb started steps=%s changed_layer_count=%s changed_slot_count=%s", self.steps, - result["before"]["max"], - result["after"]["max"], - result["before"]["p95"], - result["after"]["p95"], - result["rebalance_gain"], - result["changed_layer_count"], + sum(current != target for current, target in zip(self.current_placement, self.target_placement)), len(self.pending_transfer_infos), ) def _step_transferring(self) -> None: - """推进当前传输,并在一层完成后原子地发布该层。""" - if not self._active_transfer_finished_on_all_ranks(): + """一次启动一个专家传输,并在完成后立即提交对应槽位。""" + if not hasattr(self, "active_transfer"): + self.rebalance_started_at = time.time() + self._start_next_transfer() return - completed_info = self._complete_active_transfer() - if self._next_transfer_is_in_layer(completed_info.layer_index): - self._start_next_transfer() + if not self._active_transfer_finished_on_all_ranks(): return - self._synchronize_and_commit_layer(completed_info.layer_index) + completed_transfer = self._complete_active_transfer() + self._synchronize_and_commit_transfer(completed_transfer) if self.pending_transfer_infos: self._start_next_transfer() return @@ -290,60 +289,15 @@ def _step_transferring(self) -> None: elapsed, ) - def _publish_expert_load_metric(self, result: Dict[str, Any]) -> None: + def _publish_expert_load_metric(self, global_load: torch.Tensor) -> None: if self.global_rank != 0: return - self.metric_client.gauge_set(EPLB_EXPERT_IMBALANCE_RATIO_METRIC, result["expert_imbalance_ratio"]) - - def _build_rebalance_data(self, result: Dict[str, Any]) -> Tuple[Dict[int, torch.Tensor], List[EPLBTransferInfo]]: - metadata_by_layer: Dict[int, torch.Tensor] = {} - planned_transfers: List[EPLBTransferInfo] = [] - changed_layer_indices: List[int] = [ - layer_index for layer_index, changed in enumerate(result["changed_layers"]) if changed - ] - - num_primary_experts_per_rank = self.num_logical_experts // self.world_size - local_expert_ids_by_rank_and_layer: List[List[List[int]]] = [] - for layer_index in changed_layer_indices: - local_expert_ids_by_rank_and_layer.append( - [ - list( - range( - rank * num_primary_experts_per_rank, - (rank + 1) * num_primary_experts_per_rank, - ) - ) - + result["placement"][layer_index][rank] - for rank in range(self.world_size) - ] - ) - logical_to_physical_maps = torch.tensor( - build_logical_to_physical_maps_for_layers( - local_expert_ids_by_rank_and_layer, - self.num_logical_experts, - current_rank=self.global_rank, - ), - dtype=torch.int32, + self.metric_client.gauge_set( + EPLB_EXPERT_IMBALANCE_RATIO_METRIC, + _expert_load_imbalance_ratio(global_load), ) - for changed_layer_offset, layer_index in enumerate(changed_layer_indices): - current_layer_placement = self.current_placement[layer_index] - target_layer_placement = result["placement"][layer_index] - metadata_by_layer[layer_index] = logical_to_physical_maps[changed_layer_offset] - planned_transfers.extend( - build_transfer_plan( - current_layer_placement, - target_layer_placement, - layer_index, - self.num_logical_experts, - self.world_size, - ) - ) - return metadata_by_layer, planned_transfers - - # 传输阶段内部实现。 def _active_transfer_finished_on_all_ranks(self) -> bool: - """仅当所有 rank 都完成当前传输时返回 ``True``。""" finished_by_rank = [False] * self.world_size dist.all_gather_object( finished_by_rank, @@ -352,68 +306,70 @@ def _active_transfer_finished_on_all_ranks(self) -> bool: ) return all(finished_by_rank) - def _complete_active_transfer(self) -> EPLBTransferInfo: - """将当前传输从待处理队列移动到本层的已完成列表。""" - assert self.pending_transfer_infos - - expected_transfer_info: EPLBTransferInfo = self.pending_transfer_infos[0] + def _complete_active_transfer(self) -> PinnedMemoryEPLBTransfer: + expected_transfer_info = self.pending_transfer_infos[0] if self.active_transfer.transfer_info != expected_transfer_info: raise RuntimeError("EPLB completed transfer does not match the expected transfer info") - - self.completed_layer_transfers.append(self.active_transfer) self.pending_transfer_infos.pop(0) - return expected_transfer_info - - def _next_transfer_is_in_layer(self, layer_index: int) -> bool: - """判断下一个待处理传输是否仍属于当前层。""" - return bool(self.pending_transfer_infos and self.pending_transfer_infos[0].layer_index == layer_index) + return self.active_transfer def _start_next_transfer(self) -> None: - transfer_info: EPLBTransferInfo = self.pending_transfer_infos[0] self.active_transfer = PinnedMemoryEPLBTransfer( self._weights, self.transfer_group, self.global_rank, - transfer_info, + self.pending_transfer_infos[0], ) self.active_transfer.start() - def _synchronize_and_commit_layer(self, layer_index: int) -> None: - """等待旧权重使用完毕,然后发布一层的新权重和 metadata。""" + def _synchronize_and_commit_transfer(self, transfer: PinnedMemoryEPLBTransfer) -> None: from lightllm.server.router.model_infer.infer_batch import g_infer_context torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - self._commit_transferred_layer(layer_index) - self.completed_layer_transfers.clear() + transfer_info = transfer.transfer_info + if transfer_info.dest_rank == self.global_rank: + for tensor_buffer in transfer.tensor_buffers: + tensor_buffer.live_tensor[transfer_info.dest_local_expert_index].copy_(tensor_buffer.pinned_row) + + layer_index = transfer_info.layer_index + redundant_slot_index = transfer_info.dest_local_expert_index - self.num_primary_experts_per_rank + self.current_placement[layer_index][transfer_info.dest_rank][ + redundant_slot_index + ] = transfer_info.source_logical_expert_id + if transfer_info.dest_rank == self.global_rank: + self._eplb_impls[layer_index].local_logics_expert_ids_list[ + transfer_info.dest_local_expert_index + ] = transfer_info.source_logical_expert_id + + local_expert_ids_by_rank = [ + list( + range( + rank * self.num_primary_experts_per_rank, + (rank + 1) * self.num_primary_experts_per_rank, + ) + ) + + self.current_placement[layer_index][rank] + for rank in range(self.world_size) + ] + logical_to_physical_map = torch.tensor( + build_logical_to_physical_map( + local_expert_ids_by_rank, + self.num_logical_experts, + current_rank=self.global_rank, + ), + dtype=torch.int32, + ) + self._eplb_impls[layer_index].logical_to_physical_map.copy_(logical_to_physical_map) def _complete_rebalance(self) -> float: self.current_placement = self.target_placement elapsed = time.time() - self.rebalance_started_at del self.pending_transfer_infos - del self.completed_layer_transfers del self.active_transfer del self.target_placement - del self.target_metadata del self.rebalance_started_at return elapsed - def _commit_transferred_layer(self, layer_index: int) -> None: - """在主推理线程中同步发布一层权重和路由 metadata。""" - target_redundant_expert_ids: List[int] = self.target_placement[layer_index][self.global_rank] - for transfer in self.completed_layer_transfers: - transfer_info: EPLBTransferInfo = transfer.transfer_info - if transfer_info.dest_rank != self.global_rank: - continue - for tensor_buffer in transfer.tensor_buffers: - tensor_buffer.live_tensor[transfer_info.dest_local_expert_index].copy_(tensor_buffer.pinned_row) - - local_expert_ids: List[int] = self._eplb_impls[layer_index].local_logics_expert_ids_list - local_expert_ids[self.num_primary_experts_per_rank :] = target_redundant_expert_ids - self._commit_layer_metadata(layer_index) - - def _commit_layer_metadata(self, layer_index: int) -> None: - self._eplb_impls[layer_index].logical_to_physical_map.copy_(self.target_metadata[layer_index]) - def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: weights_by_id: Dict[int, FusedMoeWeight] = {} @@ -422,3 +378,16 @@ def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: weights_by_id[id(value)] = value return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) + + +def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: + """Average each layer's maximum-to-mean logical-expert token ratio.""" + if global_load.ndim != 2: + raise ValueError("global_load must be [layers, logical_experts]") + global_load = global_load.to(torch.float64) + layer_means = global_load.mean(dim=1) + valid_layers = layer_means > 0 + if not torch.any(valid_layers): + return 0.0 + ratios = global_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] + return float(ratios.mean().item()) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_plan.py b/lightllm/server/router/model_infer/mode_backend/eplb_plan.py index f38b3a3450..b6400a9122 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_plan.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_plan.py @@ -3,7 +3,7 @@ import os import threading from enum import Enum -from typing import Any, Dict, Optional +from typing import Optional import torch @@ -37,7 +37,7 @@ def __init__( self.global_load = global_load self.current_placement = current_placement self.status = PlanTaskStatus.IDLE - self.result: Optional[Dict[str, Any]] = None + self.result: Optional[ExpertPlacement] = None self._thread = threading.Thread( target=self._run, name="eplb-plan", @@ -56,26 +56,11 @@ def is_finished(self) -> bool: def _run(self) -> None: try: - result = self.planner.plan( + self.result = self.planner.plan( self.global_load.tolist(), self.current_placement, - ).as_dict() - result["expert_imbalance_ratio"] = _expert_load_imbalance_ratio(self.global_load) - self.result = result + ) self.status = PlanTaskStatus.SUCCEEDED except BaseException: logger.exception("EPLB planning failed") os._exit(1) - - -def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: - """Average each layer's maximum-to-mean logical-expert token ratio.""" - if global_load.ndim != 2: - raise ValueError("global_load must be [layers, logical_experts]") - global_load = global_load.to(torch.float64) - layer_means = global_load.mean(dim=1) - valid_layers = layer_means > 0 - if not torch.any(valid_layers): - return 0.0 - ratios = global_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] - return float(ratios.mean().item()) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index a694e01076..adf34d0acc 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -306,14 +306,15 @@ def test_eplb_planner_builds_legal_concrete_slot_layout(): load[:, :, 4] = 500 result = planner.plan(load.sum(dim=1).tolist(), current) - placement = result.placement[0] + placement = result[0] for rank, row in enumerate(placement): assert len(row) == len(set(row)) assert all(expert // 2 != rank for expert in row) - assert max(map(max, result.after_rank_load)) <= max(map(max, result.before_rank_load)) - assert isinstance(result.placement, list) - assert isinstance(result.changed_layers, list) + assert max(map(max, planner.estimate_rank_load(load.sum(dim=1).tolist(), result))) <= max( + map(max, planner.estimate_rank_load(load.sum(dim=1).tolist(), current)) + ) + assert isinstance(result, list) def test_eplb_planner_estimator_distributes_global_load_across_copies(): @@ -336,8 +337,7 @@ def test_eplb_planner_does_not_move_zero_load_experts(): result = planner.plan([[0, 0, 0, 0]], current) - assert result.reason == "no_improvement" - assert result.placement == current + assert result == current def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): @@ -356,9 +356,9 @@ def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): result = planner.plan(load.sum(dim=1).tolist(), current) - assert len(result.placement) == len(current) - assert all(len(actual) == len(expected) for actual, expected in zip(result.placement[0], current[0])) - for rank, row in enumerate(result.placement[0]): + assert len(result) == len(current) + assert all(len(actual) == len(expected) for actual, expected in zip(result[0], current[0])) + for rank, row in enumerate(result[0]): assert all(expert // 4 != rank for expert in row) @@ -371,7 +371,7 @@ def test_eplb_planner_keeps_high_redundancy_search_state_isolated(): result = planner.plan(load, current) - for rank, row in enumerate(result.placement[0]): + for rank, row in enumerate(result[0]): assert len(row) == len(set(row)) == 3 assert all(expert // 4 != rank for expert in row) @@ -591,18 +591,14 @@ def test_manager_delegates_distribution_planning_to_planner_class(): current_placement = [[[1]]] logical_load = torch.tensor([[10, 20]]) calls = [] - - class Result: - def as_dict(self): - return {"kind": "no_improvement"} - - planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or Result())) + planned_placement = [[[0]]] + planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or planned_placement)) task = plan_module.EPLBPlanTask(planner, logical_load, current_placement) task._run() assert task.status is plan_module.PlanTaskStatus.SUCCEEDED - assert task.result == {"kind": "no_improvement", "expert_imbalance_ratio": 4 / 3} + assert task.result == planned_placement assert len(calls) == 1 assert calls[0][0] == logical_load.tolist() assert calls[0][1] == current_placement @@ -637,7 +633,7 @@ def test_expert_load_imbalance_ratio_averages_layer_ratios(): dtype=torch.int64, ) - ratio = plan_module._expert_load_imbalance_ratio(global_load) + ratio = manager_module._expert_load_imbalance_ratio(global_load) assert ratio == pytest.approx(1.25) @@ -648,11 +644,7 @@ def test_manager_publishes_expert_load_metrics_from_rank_zero(): manager.global_rank = 0 manager.metric_client = SimpleNamespace(gauge_set=lambda name, value: calls.append((name, value))) - manager._publish_expert_load_metric( - { - "expert_imbalance_ratio": 1.25, - } - ) + manager._publish_expert_load_metric(torch.tensor([[2, 4, 6], [10, 10, 10]])) assert calls == [ (manager_module.EPLB_EXPERT_IMBALANCE_RATIO_METRIC, 1.25), @@ -663,7 +655,7 @@ def test_manager_does_not_publish_expert_load_metrics_from_other_ranks(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.global_rank = 1 - manager._publish_expert_load_metric({"expert_imbalance_ratio": 1.25}) + manager._publish_expert_load_metric(torch.tensor([[1, 2]])) assert not hasattr(manager, "metric_client") @@ -757,76 +749,6 @@ def all_gather_object(output, local_token_count, **kwargs): assert manager.state is manager_module.EPLBManagerState.COLLECTING -def test_manager_planning_builds_improved_metadata_in_one_multilayer_call( - monkeypatch, -): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "fuse_moe_impl": _test_moe_impl( - eplb=True, - route_counter=torch.zeros((4,), dtype=torch.int64), - num_logical_experts=4, - world_size=4, - ) - }, - )() - for _ in range(3) - ] - manager._eplb_impls = [weight.fuse_moe_impl for weight in manager.weights] - manager._weights = manager.weights - for layer_num, weight in enumerate(manager._weights): - weight.layer_num_ = layer_num - manager.global_rank = 1 - manager.world_size = 4 - manager.step_interval = 20 - manager.num_logical_experts = 4 - manager.num_redundant_experts_per_rank = 1 - manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone().tolist() - planned_placement = torch.tensor( - [ - [[3], [0], [1], [2]], - [[2], [3], [0], [1]], - [[2], [3], [0], [1]], - ], - dtype=torch.int64, - ) - result = { - "kind": "planned", - "placement": planned_placement.tolist(), - "changed_layers": [True, False, True], - } - calls = [] - original_build_maps_for_layers = manager_module.build_logical_to_physical_maps_for_layers - - def build_maps_for_layers(*args, **kwargs): - placements = args[0] - calls.append((len(placements), len(placements[0]), len(placements[0][0]))) - return original_build_maps_for_layers(*args, **kwargs) - - monkeypatch.setattr( - manager_module, - "build_logical_to_physical_maps_for_layers", - build_maps_for_layers, - ) - metadata, transfer_infos = manager._build_rebalance_data(result) - - assert calls == [(2, 4, 2)] - assert 1 not in metadata - assert {info.layer_index for info in transfer_infos} == {0, 2} - for layer_index in (0, 2): - item = metadata[layer_index] - expected = build_logical_to_physical_map( - _rank_to_logic_expert_ids(planned_placement[layer_index].tolist(), 4), - 4, - current_rank=manager.global_rank, - ) - assert torch.equal(item, torch.tensor(expected, dtype=torch.int32)) - - def test_decode_dispatch_uses_physical_ids_and_total_expert_count(monkeypatch): class Buffer: def low_latency_dispatch(self, **kwargs): @@ -1226,73 +1148,97 @@ def __init__(self, offset, scale=True, zero_point=True): ] -def test_manager_commits_completed_transfer_rows_by_planned_destination_index(): +def test_manager_commits_each_transfer_without_aggregation(monkeypatch): live = torch.arange(20).reshape(5, 4) original_primary = live[:3].clone() - local_expert_ids = [0, 1, 2, 4, 5] + local_expert_ids = [0, 1, 2, 3, 2] + target_placement = [[[4, 5], [0, 2]]] + expected_metadata = torch.tensor( + build_logical_to_physical_map( + _rank_to_logic_expert_ids(target_placement[0], 6), + 6, + current_rank=0, + ), + dtype=torch.int32, + ) + logical_to_physical_map = torch.zeros_like(expected_metadata) manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.global_rank = 0 + manager.world_size = 2 + manager.num_logical_experts = 6 manager.num_primary_experts_per_rank = 3 - manager.target_placement = [[[7, 6]]] - manager._eplb_impls = [SimpleNamespace(local_logics_expert_ids_list=local_expert_ids)] - manager._commit_layer_metadata = lambda _layer: None - manager.completed_layer_transfers = [ + manager.target_placement = target_placement + manager.current_placement = [[[3, 2], [0, 2]]] + manager._eplb_impls = [ SimpleNamespace( - transfer_info=EPLBTransferInfo(1, 0, 7, 0, 3), - tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -7))], + local_logics_expert_ids_list=local_expert_ids, + logical_to_physical_map=logical_to_physical_map, + ) + ] + transfers = [ + SimpleNamespace( + transfer_info=EPLBTransferInfo(1, 0, 4, 0, 3), + tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -4))], ), SimpleNamespace( - transfer_info=EPLBTransferInfo(2, 0, 6, 0, 4), - tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -6))], + transfer_info=EPLBTransferInfo(1, 0, 5, 0, 4), + tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -5))], ), ] + waits = [] + + class CurrentStream: + def wait_stream(self, stream): + waits.append(stream) - manager._commit_transferred_layer(0) + overlap_stream = object() + monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: CurrentStream()) + monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) + + manager._synchronize_and_commit_transfer(transfers[0]) + assert manager.current_placement[0][0] == [4, 2] + manager._synchronize_and_commit_transfer(transfers[1]) assert torch.equal(live[:3], original_primary) - assert torch.equal(live[3], torch.full((4,), -7)) - assert torch.equal(live[4], torch.full((4,), -6)) - assert local_expert_ids == [0, 1, 2, 7, 6] + assert torch.equal(live[3], torch.full((4,), -4)) + assert torch.equal(live[4], torch.full((4,), -5)) + assert local_expert_ids == [0, 1, 2, 4, 5] + assert torch.equal(logical_to_physical_map, expected_metadata) + assert waits == [overlap_stream, overlap_stream] -def test_manager_inflight_ready_gate_commits_one_layer(monkeypatch): +def test_manager_transfer_task_ready_gate_commits_one_layer(monkeypatch): class Transfer: def __init__(self, transfer_info): self.transfer_info = transfer_info self.status = TransferStatus.SUCCEEDED + def start(self): + starts.append(self.transfer_info) + def is_finished(self): return True info0 = EPLBTransferInfo(0, 0, 2, 1, 2) info0b = EPLBTransferInfo(1, 0, 3, 0, 2) info1 = EPLBTransferInfo(0, 1, 4, 1, 2) + starts = [] manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.active_transfer = Transfer(info0) manager.control_group = object() + manager.transfer_group = object() + manager._weights = [object(), object()] manager.world_size = 2 manager.pending_transfer_infos = [info0, info0b, info1] - manager.completed_layer_transfers = [] manager.global_rank = 1 - manager.steps = 0 - manager.step_interval = 20 - committed, finished = [], [] - manager._commit_transferred_layer = committed.append - manager._complete_rebalance = lambda: (finished.append(True) or 0.0) - - def start_next_transfer(): - manager.active_transfer = Transfer(manager.pending_transfer_infos[0]) - - manager._start_next_transfer = start_next_transfer - operations = [] - - class CurrentStream: - def wait_stream(self, stream): - operations.append(("wait", stream)) - - overlap_stream = object() - monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: CurrentStream()) - monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) + manager.state = manager_module.EPLBManagerState.TRANSFERRING + committed = [] + manager._synchronize_and_commit_transfer = lambda transfer: committed.append(transfer.transfer_info) + manager._complete_rebalance = lambda: 0.0 + monkeypatch.setattr( + manager_module, + "PinnedMemoryEPLBTransfer", + lambda _weights, _group, _rank, transfer_info: Transfer(transfer_info), + ) def set_global_ready(ready): def all_gather_object(output, _local_ready, **_kwargs): @@ -1300,23 +1246,25 @@ def all_gather_object(output, _local_ready, **_kwargs): return all_gather_object + manager._step_transferring() + assert starts == [info0] + monkeypatch.setattr(manager_module.dist, "all_gather_object", set_global_ready(False)) manager._step_transferring() assert committed == [] monkeypatch.setattr(manager_module.dist, "all_gather_object", set_global_ready(True)) manager._step_transferring() - assert committed == [] - assert operations == [] - assert not finished + assert committed == [info0] + assert starts == [info0, info0b] manager._step_transferring() - assert committed == [0] - assert operations == [("wait", overlap_stream)] + assert committed == [info0, info0b] + assert starts == [info0, info0b, info1] manager._step_transferring() - assert committed == [0, 1] - assert finished == [True] + assert committed == [info0, info0b, info1] + assert manager.state is manager_module.EPLBManagerState.COLLECTING expected = EPLBTransferInfo(0, 2, 5, 1, 2) manager.pending_transfer_infos = [expected] @@ -1346,7 +1294,7 @@ def broadcast(values, **_kwargs): @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_manager_inflight_commit_orders_live_weights_between_overlap_forwards( +def test_manager_transfer_task_commit_orders_live_weights_between_overlap_forwards( monkeypatch, ): class Transfer: @@ -1369,24 +1317,30 @@ def is_finished(self): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) transfer_info = EPLBTransferInfo(0, 0, 2, 0, 0) - manager.active_transfer = Transfer(live, received, transfer_info) + transfer = Transfer(live, received, transfer_info) + manager.active_transfer = transfer manager.control_group = object() manager.world_size = 1 manager.pending_transfer_infos = [transfer_info] - manager.completed_layer_transfers = [] manager.num_primary_experts_per_rank = 0 + manager.num_logical_experts = 1 manager.global_rank = 0 manager.target_placement = [[[2]]] - manager._eplb_impls = [SimpleNamespace(local_logics_expert_ids_list=[1])] - manager._commit_layer_metadata = lambda _layer: None - manager._complete_rebalance = lambda: 0.0 - manager.steps = 0 - manager.step_interval = 20 + manager._eplb_impls = [ + SimpleNamespace( + local_logics_expert_ids_list=[1], + logical_to_physical_map=torch.zeros((1, 1), dtype=torch.int32, device="cuda"), + ) + ] + manager.current_placement = [[[1]]] + manager.rebalance_started_at = time.time() + manager.state = manager_module.EPLBManagerState.TRANSFERRING monkeypatch.setattr( manager_module.dist, "all_gather_object", lambda output, local_ready, **_kwargs: output.__setitem__(slice(None), [local_ready]), ) + monkeypatch.setattr(manager_module, "build_logical_to_physical_map", lambda *_args, **_kwargs: [[0]]) try: g_infer_context.overlap_stream = source_stream @@ -1458,7 +1412,7 @@ def all_gather_object(output, local_token_count, **_kwargs): def test_manager_step_uses_explicit_state_instead_of_pending_work(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.TRANSFERRING - manager.pending_transfer_infos = [EPLBTransferInfo(0, 0, 2, 1, 2)] + manager.pending_transfer_infos = [object()] manager._plan_task = object() calls = [] manager._step_transferring = lambda: calls.append("transfer") @@ -1471,31 +1425,47 @@ def test_manager_step_uses_explicit_state_instead_of_pending_work(): def test_manager_enters_transferring_state_with_planned_work(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) - starts = [] manager.state = manager_module.EPLBManagerState.WAIT_PLAN_FINISH manager.global_rank = 1 manager.control_group = object() - result = { - "kind": "planned", - "placement": [[[2], [3]]], - "expert_imbalance_ratio": 1.0, - } + manager.transfer_group = object() + manager.world_size = 2 + manager.num_logical_experts = 4 + placement = [[[2], [3]], [[0], [1]]] + manager.current_placement = [[[3], [2]], [[1], [0]]] monkeypatch.setattr( manager_module.dist, "broadcast_object_list", - lambda values, **_kwargs: values.__setitem__(0, result), + lambda values, **_kwargs: values.__setitem__(0, placement), + ) + + transfer_infos = [ + EPLBTransferInfo(0, 0, 2, 1, 2), + EPLBTransferInfo(0, 1, 0, 1, 2), + ] + build_calls = [] + + def build_plan(*args): + build_calls.append(args) + return [transfer_infos[args[2]]] + + monkeypatch.setattr(manager_module, "build_transfer_plan", build_plan) + monkeypatch.setattr( + manager_module, + "PinnedMemoryEPLBTransfer", + lambda *_args: pytest.fail("transfer object must not be built while entering TRANSFERRING"), ) - manager._publish_expert_load_metric = lambda _result: None - manager._build_rebalance_data = lambda _result: ({0: torch.tensor([1])}, [transfer_info]) - manager._start_next_transfer = lambda: starts.append(manager.pending_transfer_infos[0]) manager._step_wait_plan_finish() assert manager.state is manager_module.EPLBManagerState.TRANSFERRING - assert manager.pending_transfer_infos == [transfer_info] - assert manager.completed_layer_transfers == [] - assert starts == [transfer_info] + assert manager.pending_transfer_infos == transfer_infos + assert build_calls == [ + (manager.current_placement[0], placement[0], 0, manager.num_logical_experts, manager.world_size), + (manager.current_placement[1], placement[1], 1, manager.num_logical_experts, manager.world_size), + ] + assert not hasattr(manager, "active_transfer") + assert not hasattr(manager, "target_metadata") def test_manager_evaluation_with_insufficient_tokens_returns_to_collecting(monkeypatch): @@ -1555,6 +1525,7 @@ def start(self): manager.planner = object() manager.current_placement = [[[1]]] + manager._publish_expert_load_metric = lambda _global_load: None monkeypatch.setattr(manager_module, "EPLBPlanTask", PlanTask) manager.step() @@ -1604,7 +1575,7 @@ def test_manager_planning_without_changes_returns_to_collecting(monkeypatch): manager.steps = 11 manager.step_interval = 20 manager.next_evaluation_step = 31 - manager._publish_expert_load_metric = lambda _result: None + manager.current_placement = [[[1], [0]]] result = None def broadcast(values, **_kwargs): @@ -1615,7 +1586,7 @@ def broadcast(values, **_kwargs): manager.step() assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH - result = {"kind": "no_improvement", "expert_imbalance_ratio": 1.0} + result = [[[1], [0]]] manager.step() assert manager.state is manager_module.EPLBManagerState.COLLECTING @@ -1623,32 +1594,21 @@ def broadcast(values, **_kwargs): assert not hasattr(manager, "_plan_task") -def test_manager_complete_rebalance_releases_transferring_state(): +def test_manager_complete_rebalance_releases_transfer_info_list(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) target_placement = [[[2], [3]]] manager.global_rank = 1 - manager.steps = 10 - manager.step_interval = 20 manager.pending_transfer_infos = [] - manager.completed_layer_transfers = [] manager.active_transfer = object() manager.target_placement = target_placement - manager.target_metadata = {} manager.rebalance_started_at = time.time() elapsed = manager._complete_rebalance() assert manager.current_placement is target_placement assert elapsed >= 0 - for attribute in ( - "pending_transfer_infos", - "completed_layer_transfers", - "active_transfer", - "target_placement", - "target_metadata", - "rebalance_started_at", - ): - assert not hasattr(manager, attribute) + assert not hasattr(manager, "pending_transfer_infos") + assert not hasattr(manager, "active_transfer") def test_manager_exposes_one_lifecycle_step_entrypoint(): @@ -1821,7 +1781,7 @@ def test_manager_requires_more_than_one_rank(monkeypatch): manager_module.EPLBManager(type("Model", (), {})()) -def test_manager_constructs_pinned_memory_transfer(monkeypatch): +def test_manager_initializes_without_transfer_task(monkeypatch): weight = type( "Weight", (), @@ -1838,8 +1798,6 @@ def test_manager_constructs_pinned_memory_transfer(monkeypatch): ), }, )() - transfer_starts = [] - transfer = SimpleNamespace(start=lambda: transfer_starts.append(True)) groups = [object(), object()] new_group_calls = [] monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) @@ -1868,27 +1826,15 @@ def all_gather_object(output, local_redundant_expert_ids_by_layer, group): output[:] = [local_redundant_expert_ids_by_layer, [[0, 1]]] monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) - transfer_calls = [] - monkeypatch.setattr( - manager_module, - "PinnedMemoryEPLBTransfer", - lambda weights, group, rank, info: (transfer_calls.append((weights, group, rank, info)) or transfer), - ) logs = [] monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) manager = manager_module.EPLBManager(type("Model", (), {})()) assert not hasattr(manager, "_plan_task") - assert not hasattr(manager, "active_transfer") - assert not hasattr(manager, "target_placement") + assert not hasattr(manager, "pending_transfer_infos") assert manager.state is manager_module.EPLBManagerState.COLLECTING assert (manager.control_group, manager.transfer_group) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 2 assert all_gather_calls == [([[2, 3]], groups[0])] - transfer_info = EPLBTransferInfo(0, 0, 2, 1, 2) - manager.pending_transfer_infos = [transfer_info] - manager._start_next_transfer() - assert transfer_calls == [([weight], groups[1], 0, transfer_info)] - assert transfer_starts == [True] assert manager.planner.rebalance_gain_threshold == 0.07 assert manager.current_placement == [[[2, 3], [0, 1]]] assert manager.metric_client is metric_client From facf59c7a5bee25049ef54c0e6cd49c18559f8a7 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 13:02:26 +0000 Subject: [PATCH 37/72] refine EPLB batched transfer scheduling --- .../model_infer/mode_backend/eplb_manager.py | 158 +++++++++++------- unit_tests/common/fused_moe/test_eplb.py | 120 ++++++------- 2 files changed, 157 insertions(+), 121 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index c36e814d7b..026324e0d7 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -1,6 +1,6 @@ from enum import Enum import time -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Tuple import torch import torch.distributed as dist @@ -266,69 +266,104 @@ def _step_wait_plan_finish(self) -> None: ) def _step_transferring(self) -> None: - """一次启动一个专家传输,并在完成后立即提交对应槽位。""" - if not hasattr(self, "active_transfer"): + """启动或轮询一个传输批次;整批完成后再统一提交。""" + if not hasattr(self, "rebalance_started_at"): self.rebalance_started_at = time.time() - self._start_next_transfer() - return - - if not self._active_transfer_finished_on_all_ranks(): - return - completed_transfer = self._complete_active_transfer() - self._synchronize_and_commit_transfer(completed_transfer) - if self.pending_transfer_infos: - self._start_next_transfer() - return + # 没有活动批次时,所有 rank 根据相同的 pending 列表构造下一批任务。 + if not hasattr(self, "active_transfer_batch"): + self.active_transfer_batch = self._next_transfer_batch() + + # 空批次表示公共任务列表已经耗尽,所有 rank 可以同时结束重排。 + if not self.active_transfer_batch: + del self.active_transfer_batch + self.current_placement = self.target_placement + elapsed = time.time() - self.rebalance_started_at + del self.pending_transfer_infos + del self.target_placement + del self.rebalance_started_at + self.state = EPLBManagerState.COLLECTING + if self.global_rank == 0: + logger.info("eplb completed wall_time=%.2fs", elapsed) + return + + # 一个批次内每个 rank 至多参与一条任务。参与者构造并启动本地传输; + # 其他 rank 只保存相同的批次信息,后续共同参与状态同步和 commit。 + local_transfer_info = next( + ( + transfer_info + for transfer_info in self.active_transfer_batch + if self.global_rank in (transfer_info.source_rank, transfer_info.dest_rank) + ), + None, + ) + if local_transfer_info is not None: + self.active_transfer = PinnedMemoryEPLBTransfer( + self._weights, + self.transfer_group, + self.global_rank, + local_transfer_info, + ) + self.active_transfer.start() + else: + # 已有活动批次时,本 step 只负责轮询;整批完成后才统一提交。 + self._poll_transfer_batch() + + def _next_transfer_batch(self) -> List[EPLBTransferInfo]: + """按计划顺序选择一组 rank 互不冲突的传输任务。""" + available_ranks = set(range(self.world_size)) + transfer_batch: List[EPLBTransferInfo] = [] + remaining_transfer_infos: List[EPLBTransferInfo] = [] + + # 所有 rank 持有相同的计划列表,并执行相同的贪心扫描,因此会得到完全 + # 相同的批次。每个 rank 在本批次中至多参与一条任务,可以独立启动。 + for transfer_info in self.pending_transfer_infos: + participant_ranks = {transfer_info.source_rank, transfer_info.dest_rank} + if participant_ranks <= available_ranks: + transfer_batch.append(transfer_info) + available_ranks -= participant_ranks + else: + remaining_transfer_infos.append(transfer_info) + + # 选中的任务由 active_transfer_batch 持有;未选中的冲突任务保留到下一批。 + self.pending_transfer_infos = remaining_transfer_infos + return transfer_batch + + def _poll_transfer_batch(self) -> None: + """等待当前批次全部完成,随后统一提交并释放本地任务。""" + from lightllm.server.router.model_infer.infer_batch import g_infer_context - elapsed = self._complete_rebalance() - self.state = EPLBManagerState.COLLECTING - if self.global_rank == 0: - logger.info( - "eplb completed wall_time=%.2fs", - elapsed, + local_state: Optional[Tuple[EPLBTransferInfo, bool]] = None + if hasattr(self, "active_transfer"): + local_state = ( + self.active_transfer.transfer_info, + self.active_transfer.is_finished(), ) + transfer_states: List[Optional[Tuple[EPLBTransferInfo, bool]]] = [None] * self.world_size + dist.all_gather_object(transfer_states, local_state, group=self.control_group) - def _publish_expert_load_metric(self, global_load: torch.Tensor) -> None: - if self.global_rank != 0: + # 批次内任意任务只要缺少参与方状态,或任一参与方尚未完成,整批都不能 + # commit。下一次 step 会继续轮询同一个批次。 + if not all(state is None or state[1] for state in transfer_states): return - self.metric_client.gauge_set( - EPLB_EXPERT_IMBALANCE_RATIO_METRIC, - _expert_load_imbalance_ratio(global_load), - ) - def _active_transfer_finished_on_all_ranks(self) -> bool: - finished_by_rank = [False] * self.world_size - dist.all_gather_object( - finished_by_rank, - self.active_transfer.is_finished(), - group=self.control_group, - ) - return all(finished_by_rank) - - def _complete_active_transfer(self) -> PinnedMemoryEPLBTransfer: - expected_transfer_info = self.pending_transfer_infos[0] - if self.active_transfer.transfer_info != expected_transfer_info: - raise RuntimeError("EPLB completed transfer does not match the expected transfer info") - self.pending_transfer_infos.pop(0) - return self.active_transfer - - def _start_next_transfer(self) -> None: - self.active_transfer = PinnedMemoryEPLBTransfer( - self._weights, - self.transfer_group, - self.global_rank, - self.pending_transfer_infos[0], - ) - self.active_transfer.start() + # 所有 rank 使用相同的批次顺序提交,因此全局 placement 和 metadata + # 始终一致;只有 destination rank 会额外写入实际专家权重。 + torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) + for transfer_info in self.active_transfer_batch: + self._commit_transfer(transfer_info) - def _synchronize_and_commit_transfer(self, transfer: PinnedMemoryEPLBTransfer) -> None: - from lightllm.server.router.model_infer.infer_batch import g_infer_context + if hasattr(self, "active_transfer"): + del self.active_transfer + del self.active_transfer_batch - torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - transfer_info = transfer.transfer_info + def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: + """提交一条传输,并发布更新后的路由 metadata。""" if transfer_info.dest_rank == self.global_rank: - for tensor_buffer in transfer.tensor_buffers: + assert ( + hasattr(self, "active_transfer") and self.active_transfer.transfer_info == transfer_info + ), "EPLB destination rank has no matching completed transfer" + for tensor_buffer in self.active_transfer.tensor_buffers: tensor_buffer.live_tensor[transfer_info.dest_local_expert_index].copy_(tensor_buffer.pinned_row) layer_index = transfer_info.layer_index @@ -361,14 +396,13 @@ def _synchronize_and_commit_transfer(self, transfer: PinnedMemoryEPLBTransfer) - ) self._eplb_impls[layer_index].logical_to_physical_map.copy_(logical_to_physical_map) - def _complete_rebalance(self) -> float: - self.current_placement = self.target_placement - elapsed = time.time() - self.rebalance_started_at - del self.pending_transfer_infos - del self.active_transfer - del self.target_placement - del self.rebalance_started_at - return elapsed + def _publish_expert_load_metric(self, global_load: torch.Tensor) -> None: + if self.global_rank != 0: + return + self.metric_client.gauge_set( + EPLB_EXPERT_IMBALANCE_RATIO_METRIC, + _expert_load_imbalance_ratio(global_load), + ) def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index adf34d0acc..fb4350d46a 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1148,7 +1148,7 @@ def __init__(self, offset, scale=True, zero_point=True): ] -def test_manager_commits_each_transfer_without_aggregation(monkeypatch): +def test_manager_commits_transfer_rows_and_metadata(): live = torch.arange(20).reshape(5, 4) original_primary = live[:3].clone() local_expert_ids = [0, 1, 2, 3, 2] @@ -1185,92 +1185,110 @@ def test_manager_commits_each_transfer_without_aggregation(monkeypatch): tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -5))], ), ] - waits = [] - - class CurrentStream: - def wait_stream(self, stream): - waits.append(stream) - - overlap_stream = object() - monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: CurrentStream()) - monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) - - manager._synchronize_and_commit_transfer(transfers[0]) + manager.active_transfer = transfers[0] + manager._commit_transfer(transfers[0].transfer_info) assert manager.current_placement[0][0] == [4, 2] - manager._synchronize_and_commit_transfer(transfers[1]) + manager.active_transfer = transfers[1] + manager._commit_transfer(transfers[1].transfer_info) assert torch.equal(live[:3], original_primary) assert torch.equal(live[3], torch.full((4,), -4)) assert torch.equal(live[4], torch.full((4,), -5)) assert local_expert_ids == [0, 1, 2, 4, 5] assert torch.equal(logical_to_physical_map, expected_metadata) - assert waits == [overlap_stream, overlap_stream] -def test_manager_transfer_task_ready_gate_commits_one_layer(monkeypatch): +def test_manager_transfers_only_local_tasks_and_gathers_global_status(monkeypatch): class Transfer: def __init__(self, transfer_info): self.transfer_info = transfer_info - self.status = TransferStatus.SUCCEEDED + self.finished = finished_by_info[transfer_info] def start(self): starts.append(self.transfer_info) def is_finished(self): - return True + return self.finished - info0 = EPLBTransferInfo(0, 0, 2, 1, 2) - info0b = EPLBTransferInfo(1, 0, 3, 0, 2) - info1 = EPLBTransferInfo(0, 1, 4, 1, 2) + remote_info = EPLBTransferInfo(0, 0, 2, 2, 2) + local_info0 = EPLBTransferInfo(1, 0, 3, 3, 2) + local_info1 = EPLBTransferInfo(0, 1, 4, 1, 2) + finished_by_info = {local_info0: False, local_info1: True} starts = [] manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.control_group = object() manager.transfer_group = object() manager._weights = [object(), object()] - manager.world_size = 2 - manager.pending_transfer_infos = [info0, info0b, info1] + manager.world_size = 4 + manager.pending_transfer_infos = [remote_info, local_info0, local_info1] + manager.target_placement = [[[2]], [[4]]] + manager.current_placement = [[[0]], [[3]]] manager.global_rank = 1 manager.state = manager_module.EPLBManagerState.TRANSFERRING committed = [] - manager._synchronize_and_commit_transfer = lambda transfer: committed.append(transfer.transfer_info) - manager._complete_rebalance = lambda: 0.0 + manager._commit_transfer = committed.append + waits = [] + overlap_stream = object() + monkeypatch.setattr( + manager_module.torch.cuda, + "current_stream", + lambda: SimpleNamespace(wait_stream=lambda stream: waits.append(stream)), + ) + monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) monkeypatch.setattr( manager_module, "PinnedMemoryEPLBTransfer", lambda _weights, _group, _rank, transfer_info: Transfer(transfer_info), ) - def set_global_ready(ready): - def all_gather_object(output, _local_ready, **_kwargs): - output[:] = [ready] * manager.world_size + def state(transfer_info, finished): + return transfer_info, finished + + gathered_states = [ + [state(remote_info, True), state(local_info0, False), state(remote_info, True), state(local_info0, False)], + [state(remote_info, True), state(local_info0, True), state(remote_info, True), state(local_info0, True)], + [state(local_info1, True), state(local_info1, True), None, None], + ] + local_states = [] + + def all_gather_object(output, local_state, **_kwargs): + local_states.append(local_state) + output[:] = gathered_states.pop(0) - return all_gather_object + monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) manager._step_transferring() - assert starts == [info0] + assert starts == [local_info0] + assert committed == [] + assert manager.active_transfer_batch == [remote_info, local_info0] + assert manager.pending_transfer_infos == [local_info1] - monkeypatch.setattr(manager_module.dist, "all_gather_object", set_global_ready(False)) manager._step_transferring() assert committed == [] - monkeypatch.setattr(manager_module.dist, "all_gather_object", set_global_ready(True)) + manager.active_transfer.finished = True manager._step_transferring() - assert committed == [info0] - assert starts == [info0, info0b] + assert committed == [remote_info, local_info0] + assert not hasattr(manager, "active_transfer") manager._step_transferring() - assert committed == [info0, info0b] - assert starts == [info0, info0b, info1] + assert starts == [local_info0, local_info1] manager._step_transferring() - assert committed == [info0, info0b, info1] - assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert committed == [remote_info, local_info0, local_info1] - expected = EPLBTransferInfo(0, 2, 5, 1, 2) - manager.pending_transfer_infos = [expected] - manager.active_transfer = Transfer(EPLBTransferInfo(1, 2, 5, 1, 2)) - with pytest.raises(RuntimeError, match="does not match"): - manager._step_transferring() + manager._step_transferring() + assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert manager.current_placement == [[[2]], [[4]]] + assert not hasattr(manager, "pending_transfer_infos") + assert not hasattr(manager, "target_placement") + assert not hasattr(manager, "rebalance_started_at") + assert local_states == [ + state(local_info0, False), + state(local_info0, True), + state(local_info1, True), + ] + assert waits == [overlap_stream, overlap_stream] def test_wait_plan_finish_broadcasts_pending_status(monkeypatch): @@ -1319,6 +1337,7 @@ def is_finished(self): transfer_info = EPLBTransferInfo(0, 0, 2, 0, 0) transfer = Transfer(live, received, transfer_info) manager.active_transfer = transfer + manager.active_transfer_batch = [transfer_info] manager.control_group = object() manager.world_size = 1 manager.pending_transfer_infos = [transfer_info] @@ -1594,23 +1613,6 @@ def broadcast(values, **_kwargs): assert not hasattr(manager, "_plan_task") -def test_manager_complete_rebalance_releases_transfer_info_list(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - target_placement = [[[2], [3]]] - manager.global_rank = 1 - manager.pending_transfer_infos = [] - manager.active_transfer = object() - manager.target_placement = target_placement - manager.rebalance_started_at = time.time() - - elapsed = manager._complete_rebalance() - - assert manager.current_placement is target_placement - assert elapsed >= 0 - assert not hasattr(manager, "pending_transfer_infos") - assert not hasattr(manager, "active_transfer") - - def test_manager_exposes_one_lifecycle_step_entrypoint(): assert hasattr(manager_module.EPLBManager, "step") assert not hasattr(manager_module.EPLBManager, "poll") From 5b52214e3df191a2eeb9135e205d4f4e1b3905e5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 13:14:07 +0000 Subject: [PATCH 38/72] simplify EPLB transfer batch flow --- .../model_infer/mode_backend/eplb_manager.py | 53 +++++++++++-------- 1 file changed, 31 insertions(+), 22 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 026324e0d7..0084afcffe 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -258,10 +258,13 @@ def _step_wait_plan_finish(self) -> None: raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") self.state = EPLBManagerState.TRANSFERRING if self.global_rank == 0: + changed_layer_count = sum( + current != target for current, target in zip(self.current_placement, self.target_placement) + ) logger.info( "eplb started steps=%s changed_layer_count=%s changed_slot_count=%s", self.steps, - sum(current != target for current, target in zip(self.current_placement, self.target_placement)), + changed_layer_count, len(self.pending_transfer_infos), ) @@ -272,11 +275,10 @@ def _step_transferring(self) -> None: # 没有活动批次时,所有 rank 根据相同的 pending 列表构造下一批任务。 if not hasattr(self, "active_transfer_batch"): - self.active_transfer_batch = self._next_transfer_batch() + transfer_batch = self._pop_next_transfer_batch() # 空批次表示公共任务列表已经耗尽,所有 rank 可以同时结束重排。 - if not self.active_transfer_batch: - del self.active_transfer_batch + if not transfer_batch: self.current_placement = self.target_placement elapsed = time.time() - self.rebalance_started_at del self.pending_transfer_infos @@ -287,12 +289,14 @@ def _step_transferring(self) -> None: logger.info("eplb completed wall_time=%.2fs", elapsed) return + self.active_transfer_batch = transfer_batch + # 一个批次内每个 rank 至多参与一条任务。参与者构造并启动本地传输; # 其他 rank 只保存相同的批次信息,后续共同参与状态同步和 commit。 local_transfer_info = next( ( transfer_info - for transfer_info in self.active_transfer_batch + for transfer_info in transfer_batch if self.global_rank in (transfer_info.source_rank, transfer_info.dest_rank) ), None, @@ -309,9 +313,9 @@ def _step_transferring(self) -> None: # 已有活动批次时,本 step 只负责轮询;整批完成后才统一提交。 self._poll_transfer_batch() - def _next_transfer_batch(self) -> List[EPLBTransferInfo]: - """按计划顺序选择一组 rank 互不冲突的传输任务。""" - available_ranks = set(range(self.world_size)) + def _pop_next_transfer_batch(self) -> List[EPLBTransferInfo]: + """从 pending 中取出一组 rank 互不冲突的传输任务。""" + occupied_ranks = set() transfer_batch: List[EPLBTransferInfo] = [] remaining_transfer_infos: List[EPLBTransferInfo] = [] @@ -319,9 +323,10 @@ def _next_transfer_batch(self) -> List[EPLBTransferInfo]: # 相同的批次。每个 rank 在本批次中至多参与一条任务,可以独立启动。 for transfer_info in self.pending_transfer_infos: participant_ranks = {transfer_info.source_rank, transfer_info.dest_rank} - if participant_ranks <= available_ranks: + # 与已选任务没有公共 rank 时,当前任务可以并入本批次。 + if not participant_ranks & occupied_ranks: transfer_batch.append(transfer_info) - available_ranks -= participant_ranks + occupied_ranks |= participant_ranks else: remaining_transfer_infos.append(transfer_info) @@ -331,13 +336,12 @@ def _next_transfer_batch(self) -> List[EPLBTransferInfo]: def _poll_transfer_batch(self) -> None: """等待当前批次全部完成,随后统一提交并释放本地任务。""" - from lightllm.server.router.model_infer.infer_batch import g_infer_context - + active_transfer = getattr(self, "active_transfer", None) local_state: Optional[Tuple[EPLBTransferInfo, bool]] = None - if hasattr(self, "active_transfer"): + if active_transfer is not None: local_state = ( - self.active_transfer.transfer_info, - self.active_transfer.is_finished(), + active_transfer.transfer_info, + active_transfer.is_finished(), ) transfer_states: List[Optional[Tuple[EPLBTransferInfo, bool]]] = [None] * self.world_size dist.all_gather_object(transfer_states, local_state, group=self.control_group) @@ -349,30 +353,35 @@ def _poll_transfer_batch(self) -> None: # 所有 rank 使用相同的批次顺序提交,因此全局 placement 和 metadata # 始终一致;只有 destination rank 会额外写入实际专家权重。 + from lightllm.server.router.model_infer.infer_batch import g_infer_context + torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) for transfer_info in self.active_transfer_batch: self._commit_transfer(transfer_info) - if hasattr(self, "active_transfer"): + if active_transfer is not None: del self.active_transfer del self.active_transfer_batch def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: """提交一条传输,并发布更新后的路由 metadata。""" - if transfer_info.dest_rank == self.global_rank: + is_destination_rank = transfer_info.dest_rank == self.global_rank + if is_destination_rank: + active_transfer = getattr(self, "active_transfer", None) assert ( - hasattr(self, "active_transfer") and self.active_transfer.transfer_info == transfer_info + active_transfer is not None and active_transfer.transfer_info == transfer_info ), "EPLB destination rank has no matching completed transfer" - for tensor_buffer in self.active_transfer.tensor_buffers: + for tensor_buffer in active_transfer.tensor_buffers: tensor_buffer.live_tensor[transfer_info.dest_local_expert_index].copy_(tensor_buffer.pinned_row) layer_index = transfer_info.layer_index + layer_impl = self._eplb_impls[layer_index] redundant_slot_index = transfer_info.dest_local_expert_index - self.num_primary_experts_per_rank self.current_placement[layer_index][transfer_info.dest_rank][ redundant_slot_index ] = transfer_info.source_logical_expert_id - if transfer_info.dest_rank == self.global_rank: - self._eplb_impls[layer_index].local_logics_expert_ids_list[ + if is_destination_rank: + layer_impl.local_logics_expert_ids_list[ transfer_info.dest_local_expert_index ] = transfer_info.source_logical_expert_id @@ -394,7 +403,7 @@ def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: ), dtype=torch.int32, ) - self._eplb_impls[layer_index].logical_to_physical_map.copy_(logical_to_physical_map) + layer_impl.logical_to_physical_map.copy_(logical_to_physical_map) def _publish_expert_load_metric(self, global_load: torch.Tensor) -> None: if self.global_rank != 0: From 7b1091b28f4fef11d1ae226bde36534e39cd5fd5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 13:28:23 +0000 Subject: [PATCH 39/72] standardize EPLB full expert placements --- .../meta_weights/fused_moe/eplb_planner.py | 54 +++++---- .../model_infer/mode_backend/eplb_manager.py | 35 ++---- .../model_infer/mode_backend/eplb_transfer.py | 29 +++-- unit_tests/common/fused_moe/test_eplb.py | 108 +++++++++++------- .../fused_moe/test_eplb_transfer_gpu.py | 4 +- 5 files changed, 132 insertions(+), 98 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py index 4d6c4aeddd..7640a37162 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py @@ -14,7 +14,7 @@ # [layer][logical expert] LogicalExpertLoad = List[List[float]] -# [layer][rank][redundant slot] -> logical expert +# [layer][rank][local physical expert] -> logical expert ExpertPlacement = List[List[List[int]]] # [layer][rank] RankLoad = List[List[float]] @@ -29,7 +29,7 @@ def plan( logical_expert_load: LogicalExpertLoad, current_placement: ExpertPlacement, ) -> ExpertPlacement: - """Return a concrete ``[layer][rank][redundant slot]`` placement.""" + """Return a concrete ``[layer][rank][local physical expert]`` placement.""" class GreedyEPLBPlanner(EPLBPlanner): @@ -65,9 +65,9 @@ def plan( ``logical_expert_load`` is ``[layer][logical_expert]`` and contains the load summed across all ranks. - ``current_placement`` is ``[layer][rank][redundant_slot]``. The - returned placement has the same shape and already names the exact - physical slots that migration should update. + ``current_placement`` is ``[layer][rank][local_physical_expert]``. + Each rank row contains its fixed primary experts followed by its + movable redundant experts. The returned placement has the same shape. """ load = [[float(value) for value in layer] for layer in logical_expert_load] current = [[[int(expert) for expert in rank] for rank in layer] for layer in current_placement] @@ -113,8 +113,9 @@ def _plan_layer( experts_per_rank = num_logical_experts // self.world_size owner = [expert // experts_per_rank for expert in range(num_logical_experts)] copy_count = self._allocate_copy_count(logical_load, owner) + current_redundant_placement = [row[experts_per_rank:] for row in current_placement] - placement = [[-1] * self.num_redundant_experts_per_rank for _ in range(self.world_size)] + redundant_placement = [[-1] * self.num_redundant_experts_per_rank for _ in range(self.world_size)] locations = [{owner_rank} for owner_rank in owner] expert_rank_load = [ self._physical_load( @@ -141,23 +142,23 @@ def _plan_layer( for rank in range(self.world_size): if rank in locations[expert]: continue - empty_slots = [slot for slot, value in enumerate(placement[rank]) if value < 0] + empty_slots = [slot for slot, value in enumerate(redundant_placement[rank]) if value < 0] if not empty_slots: continue slot = min( empty_slots, key=lambda candidate: ( - current_placement[rank][candidate] != expert, + current_redundant_placement[rank][candidate] != expert, candidate, ), ) - trial_placement = [row[:] for row in placement] - trial_placement[rank][slot] = expert + trial_redundant_placement = [row[:] for row in redundant_placement] + trial_redundant_placement[rank][slot] = expert trial_locations = [set(ranks) for ranks in locations] trial_locations[expert].add(rank) if not self._can_complete( instances[index + 1 :], - trial_placement, + trial_redundant_placement, trial_locations, owner, ): @@ -173,7 +174,7 @@ def _plan_layer( ] candidate = ( max(trial_rank_load), - current_placement[rank][slot] != expert, + current_redundant_placement[rank][slot] != expert, sum(trial_rank_load), rank, slot, @@ -186,11 +187,14 @@ def _plan_layer( if best is None: raise RuntimeError("EPLB planner found no valid redundant expert placement") _, _, _, rank, slot, rank_load, next_expert_load = best - placement[rank][slot] = expert + redundant_placement[rank][slot] = expert locations[expert].add(rank) expert_rank_load[expert] = next_expert_load - return placement + return [ + list(range(rank * experts_per_rank, (rank + 1) * experts_per_rank)) + redundant_experts + for rank, redundant_experts in enumerate(redundant_placement) + ] def _allocate_copy_count(self, logical_load: List[float], owner: List[int]) -> List[int]: copy_count = [1] * len(logical_load) @@ -219,13 +223,13 @@ def _allocate_copy_count(self, logical_load: List[float], owner: List[int]) -> L def _can_complete( self, remaining_instances: List[int], - placement: List[List[int]], + redundant_placement: List[List[int]], locations: List[set], owner: List[int], ) -> bool: """Check that a greedy choice leaves a legal assignment for all slots.""" remaining = Counter(remaining_instances) - capacity = [sum(expert < 0 for expert in row) for row in placement] + capacity = [sum(expert < 0 for expert in row) for row in redundant_placement] memo = set() def search() -> bool: @@ -287,8 +291,7 @@ def _expert_locations( placement: List[List[int]], num_logical_experts: int, ) -> List[set]: - experts_per_rank = num_logical_experts // self.world_size - locations = [{expert // experts_per_rank} for expert in range(num_logical_experts)] + locations = [set() for _ in range(num_logical_experts)] for rank, row in enumerate(placement): for expert in row: if rank in locations[expert]: @@ -310,6 +313,8 @@ def _validate_inputs( raise ValueError("logical expert count must be positive and divisible by world_size") if self.num_redundant_experts_per_rank > num_logical_experts - num_logical_experts // self.world_size: raise ValueError("too many redundant slots to avoid local or duplicate replicas") + num_primary_experts_per_rank = num_logical_experts // self.world_size + num_local_experts_per_rank = num_primary_experts_per_rank + self.num_redundant_experts_per_rank for layer_load, layer_placement in zip(logical_expert_load, placement): if len(layer_load) != num_logical_experts: @@ -317,10 +322,19 @@ def _validate_inputs( if any(value < 0 for value in layer_load): raise ValueError("logical expert load must be non-negative") if len(layer_placement) != self.world_size or any( - len(rank) != self.num_redundant_experts_per_rank for rank in layer_placement + len(rank) != num_local_experts_per_rank for rank in layer_placement ): - raise ValueError("each placement layer must be [world_size][redundant_slots]") + raise ValueError("each placement layer must be [world_size][local_experts]") if any(expert < 0 or expert >= num_logical_experts for rank in layer_placement for expert in rank): raise ValueError("placement contains an invalid logical expert") + for rank, local_expert_ids in enumerate(layer_placement): + expected_primary_expert_ids = list( + range( + rank * num_primary_experts_per_rank, + (rank + 1) * num_primary_experts_per_rank, + ) + ) + if local_expert_ids[:num_primary_experts_per_rank] != expected_primary_expert_ids: + raise ValueError("placement primary experts do not match their owning rank") self._expert_locations(layer_placement, num_logical_experts) return num_logical_experts diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 0084afcffe..5d584ff5fa 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -88,23 +88,21 @@ def __init__(self, model: TpPartBaseModel) -> None: self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") - # 专家布局与规划器:直接以各层 impl 中的实际冗余专家槽位为准。 - # 本 rank 的布局索引为 [layer][redundant_slot]。 - local_redundant_expert_ids_by_layer = [ - impl.local_logics_expert_ids_list[self.num_primary_experts_per_rank :] for impl in self._eplb_impls - ] + # 每层布局都保存完整的本地专家列表:固定主专家在前,冗余专家在后。 + # 本 rank 的布局索引为 [layer][local_expert]。 + local_expert_ids_by_layer = [list(impl.local_logics_expert_ids_list) for impl in self._eplb_impls] - # all_gather 后的布局索引为 [rank][layer][redundant_slot]。 - redundant_expert_ids_by_rank_and_layer: List[List[List[int]]] = [[] for _ in range(self.world_size)] + # all_gather 后的布局索引为 [rank][layer][local_expert]。 + expert_ids_by_rank_and_layer: List[List[List[int]]] = [[] for _ in range(self.world_size)] dist.all_gather_object( - redundant_expert_ids_by_rank_and_layer, - local_redundant_expert_ids_by_layer, + expert_ids_by_rank_and_layer, + local_expert_ids_by_layer, group=self.control_group, ) - # 转置为规划器使用的 [layer][rank][redundant_slot]。 + # 转置为全局统一使用的 [layer][rank][local_expert]。 self.current_placement: ExpertPlacement = [ - [redundant_expert_ids_by_rank_and_layer[rank][layer_index] for rank in range(self.world_size)] + [expert_ids_by_rank_and_layer[rank][layer_index] for rank in range(self.world_size)] for layer_index in range(len(weights)) ] self.planner: EPLBPlanner = GreedyEPLBPlanner( @@ -376,28 +374,17 @@ def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: layer_index = transfer_info.layer_index layer_impl = self._eplb_impls[layer_index] - redundant_slot_index = transfer_info.dest_local_expert_index - self.num_primary_experts_per_rank self.current_placement[layer_index][transfer_info.dest_rank][ - redundant_slot_index + transfer_info.dest_local_expert_index ] = transfer_info.source_logical_expert_id if is_destination_rank: layer_impl.local_logics_expert_ids_list[ transfer_info.dest_local_expert_index ] = transfer_info.source_logical_expert_id - local_expert_ids_by_rank = [ - list( - range( - rank * self.num_primary_experts_per_rank, - (rank + 1) * self.num_primary_experts_per_rank, - ) - ) - + self.current_placement[layer_index][rank] - for rank in range(self.world_size) - ] logical_to_physical_map = torch.tensor( build_logical_to_physical_map( - local_expert_ids_by_rank, + self.current_placement[layer_index], self.num_logical_experts, current_rank=self.global_rank, ), diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index 0793a3f989..b7498051a5 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -198,20 +198,33 @@ def build_transfer_plan( """生成一层中所有发生变化的冗余专家传输任务。 ``current_placement`` 和 ``target_placement`` 的形状均为 - ``[world_size, num_redundant_slots]``。每个逻辑专家按照无冗余布局连续分配 - 给各 rank;该固定主副本始终作为传输源,不再从已有冗余副本中选择数据源。 + ``[world_size, num_local_experts]``,每行固定主专家在前、冗余专家在后。 + 每个逻辑专家的固定主副本始终作为传输源,不从已有冗余副本中选择数据源。 """ assert world_size > 0 assert num_logical_experts % world_size == 0 assert len(current_placement) == len(target_placement) == world_size - num_redundant_slots = len(current_placement[0]) - assert all(len(row) == num_redundant_slots for row in current_placement) - assert all(len(row) == num_redundant_slots for row in target_placement) - num_primary_experts_per_rank = num_logical_experts // world_size + num_local_experts_per_rank = len(current_placement[0]) + assert num_local_experts_per_rank >= num_primary_experts_per_rank + assert all(len(row) == num_local_experts_per_rank for row in current_placement) + assert all(len(row) == num_local_experts_per_rank for row in target_placement) + + for rank, (current_row, target_row) in enumerate(zip(current_placement, target_placement)): + expected_primary_experts = list( + range( + rank * num_primary_experts_per_rank, + (rank + 1) * num_primary_experts_per_rank, + ) + ) + assert list(current_row[:num_primary_experts_per_rank]) == expected_primary_experts + assert list(target_row[:num_primary_experts_per_rank]) == expected_primary_experts + transfer_infos: List[EPLBTransferInfo] = [] for destination_rank, (current_row, target_row) in enumerate(zip(current_placement, target_placement)): - for destination_slot_index, (current_expert_id, target_expert_id) in enumerate(zip(current_row, target_row)): + for destination_local_expert_index in range(num_primary_experts_per_rank, num_local_experts_per_rank): + current_expert_id = current_row[destination_local_expert_index] + target_expert_id = target_row[destination_local_expert_index] if target_expert_id == current_expert_id: continue assert 0 <= target_expert_id < num_logical_experts @@ -222,7 +235,7 @@ def build_transfer_plan( layer_index=layer_index, source_logical_expert_id=target_expert_id, dest_rank=destination_rank, - dest_local_expert_index=num_primary_experts_per_rank + destination_slot_index, + dest_local_expert_index=destination_local_expert_index, ) ) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index fb4350d46a..bd3073fdc8 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -92,13 +92,13 @@ def _set_deepgemm_runtime(impl, runtime): setattr(impl, name, getattr(runtime, name)) -def _initial_extra_expert_placement(num_logical_experts, world_size, num_redundant_experts_per_rank): - num_primary_experts_per_rank = num_logical_experts // world_size - initial_local_expert_ids_by_rank = build_initial_local_expert_ids( - num_logical_experts, world_size, num_redundant_experts_per_rank - ) +def _initial_expert_placement(num_logical_experts, world_size, num_redundant_experts_per_rank): return torch.tensor( - [expert_ids[num_primary_experts_per_rank:] for expert_ids in initial_local_expert_ids_by_rank], + build_initial_local_expert_ids( + num_logical_experts, + world_size, + num_redundant_experts_per_rank, + ), dtype=torch.int64, ) @@ -300,7 +300,7 @@ def test_eplb_planner_builds_legal_concrete_slot_layout(): expert_alignment=1, rebalance_gain_threshold=0.0, ) - current = _initial_extra_expert_placement(8, 4, 1).unsqueeze(0).tolist() + current = _initial_expert_placement(8, 4, 1).unsqueeze(0).tolist() load = torch.ones((1, 4, 8), dtype=torch.int64) load[:, :, 0] = 1000 load[:, :, 4] = 500 @@ -309,6 +309,8 @@ def test_eplb_planner_builds_legal_concrete_slot_layout(): placement = result[0] for rank, row in enumerate(placement): + assert row[:2] == list(range(rank * 2, (rank + 1) * 2)) + row = row[2:] assert len(row) == len(set(row)) assert all(expert // 2 != rank for expert in row) assert max(map(max, planner.estimate_rank_load(load.sum(dim=1).tolist(), result))) <= max( @@ -323,7 +325,7 @@ def test_eplb_planner_estimator_distributes_global_load_across_copies(): 1, expert_alignment=128, ) - placement = [[[2], [4], [6], [0]]] + placement = [[[0, 1, 2], [2, 3, 4], [4, 5, 6], [6, 7, 0]]] load = [[100, 200, 300, 400, 500, 600, 700, 800]] predicted = planner.estimate_rank_load(load, placement) @@ -333,7 +335,7 @@ def test_eplb_planner_estimator_distributes_global_load_across_copies(): def test_eplb_planner_does_not_move_zero_load_experts(): planner = GreedyEPLBPlanner(2, 1) - current = [[[3], [1]]] + current = [[[0, 1, 3], [2, 3, 1]]] result = planner.plan([[0, 0, 0, 0]], current) @@ -346,7 +348,7 @@ def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): 1, rebalance_gain_threshold=0.0, ) - current = _initial_extra_expert_placement(16, 4, 1).unsqueeze(0).tolist() + current = _initial_expert_placement(16, 4, 1).unsqueeze(0).tolist() load = torch.randint( 0, 10000, @@ -359,12 +361,13 @@ def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): assert len(result) == len(current) assert all(len(actual) == len(expected) for actual, expected in zip(result[0], current[0])) for rank, row in enumerate(result[0]): - assert all(expert // 4 != rank for expert in row) + assert row[:4] == list(range(rank * 4, (rank + 1) * 4)) + assert all(expert // 4 != rank for expert in row[4:]) def test_eplb_planner_keeps_high_redundancy_search_state_isolated(): planner = GreedyEPLBPlanner(4, 3) - current = _initial_extra_expert_placement(16, 4, 3).unsqueeze(0).tolist() + current = _initial_expert_placement(16, 4, 3).unsqueeze(0).tolist() load = [ [22613, 26852, 21852, 23480, 13270, 14695, 28735, 22303, 15324, 19604, 21492, 25458, 14120, 12130, 18620, 22888] ] @@ -372,6 +375,8 @@ def test_eplb_planner_keeps_high_redundancy_search_state_isolated(): result = planner.plan(load, current) for rank, row in enumerate(result[0]): + assert row[:4] == list(range(rank * 4, (rank + 1) * 4)) + row = row[4:] assert len(row) == len(set(row)) == 3 assert all(expert // 4 != rank for expert in row) @@ -540,14 +545,14 @@ def test_logical_to_physical_maps_for_layers_match_single_layer_api(current_rank def test_transfer_plan_respects_explicit_target_slots(): - current = [[4, 5], [6, 7], [0, 1], [2, 3]] - target = [[5, 4], [7, 6], [1, 0], [3, 2]] + current = [[0, 1, 4, 5], [2, 3, 6, 7], [4, 5, 0, 1], [6, 7, 2, 3]] + target = [[0, 1, 5, 4], [2, 3, 7, 6], [4, 5, 1, 0], [6, 7, 3, 2]] plan = build_transfer_plan(current, target, 3, num_logical_experts=8, world_size=4) assert all(info.layer_index == 3 for info in plan) assert {(info.dest_rank, info.source_logical_expert_id) for info in plan} == { - (rank, target[rank][slot]) for rank in range(4) for slot in range(2) + (rank, target[rank][slot]) for rank in range(4) for slot in range(2, 4) } @@ -588,10 +593,10 @@ def all_gather_object(output, local_token_count, **_kwargs): def test_manager_delegates_distribution_planning_to_planner_class(): - current_placement = [[[1]]] + current_placement = [[[0, 1]]] logical_load = torch.tensor([[10, 20]]) calls = [] - planned_placement = [[[0]]] + planned_placement = [[[0, 1]]] planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or planned_placement)) task = plan_module.EPLBPlanTask(planner, logical_load, current_placement) @@ -727,7 +732,7 @@ def test_manager_evaluation_gathers_token_counts_from_all_ranks(monkeypatch): manager.step_interval = 20 manager.num_logical_experts = 4 manager.num_redundant_experts_per_rank = 1 - manager.current_placement = _initial_extra_expert_placement(4, 4, 1).unsqueeze(0).tolist() + manager.current_placement = _initial_expert_placement(4, 4, 1).unsqueeze(0).tolist() manager.control_group = object() local = torch.full((4,), 100, dtype=torch.int64) manager._eplb_impls[0].route_counter = local @@ -1096,8 +1101,8 @@ def fused(**kwargs): def test_transfer_plan_always_uses_primary_expert_rank(): - current = [[4, 5], [6, 7], [0, 1], [2, 3]] - target = [[6, 5], [6, 7], [0, 4], [2, 3]] + current = [[0, 1, 4, 5], [2, 3, 6, 7], [4, 5, 0, 1], [6, 7, 2, 3]] + target = [[0, 1, 6, 5], [2, 3, 6, 7], [4, 5, 0, 4], [6, 7, 2, 3]] plan = build_transfer_plan(current, target, 5, num_logical_experts=8, world_size=4) assert plan == [ EPLBTransferInfo(3, 5, 6, 0, 2), @@ -1106,8 +1111,8 @@ def test_transfer_plan_always_uses_primary_expert_rank(): def test_transfer_plan_uses_same_primary_source_for_repeated_expert(): - current = [[0, 1], [2, 3], [4, 5], [4, 7]] - target = [[4, 4], [2, 3], [4, 5], [4, 7]] + current = [[0, 1, 0, 1], [2, 3, 2, 3], [4, 5, 4, 5], [6, 7, 4, 7]] + target = [[0, 1, 4, 4], [2, 3, 2, 3], [4, 5, 4, 5], [6, 7, 4, 7]] first = build_transfer_plan(current, target, 5, 8, 4) second = build_transfer_plan(current, target, 5, 8, 4) assert first == second @@ -1152,10 +1157,10 @@ def test_manager_commits_transfer_rows_and_metadata(): live = torch.arange(20).reshape(5, 4) original_primary = live[:3].clone() local_expert_ids = [0, 1, 2, 3, 2] - target_placement = [[[4, 5], [0, 2]]] + target_placement = [[[0, 1, 2, 4, 5], [3, 4, 5, 0, 2]]] expected_metadata = torch.tensor( build_logical_to_physical_map( - _rank_to_logic_expert_ids(target_placement[0], 6), + target_placement[0], 6, current_rank=0, ), @@ -1168,7 +1173,7 @@ def test_manager_commits_transfer_rows_and_metadata(): manager.num_logical_experts = 6 manager.num_primary_experts_per_rank = 3 manager.target_placement = target_placement - manager.current_placement = [[[3, 2], [0, 2]]] + manager.current_placement = [[[0, 1, 2, 3, 2], [3, 4, 5, 0, 2]]] manager._eplb_impls = [ SimpleNamespace( local_logics_expert_ids_list=local_expert_ids, @@ -1187,7 +1192,7 @@ def test_manager_commits_transfer_rows_and_metadata(): ] manager.active_transfer = transfers[0] manager._commit_transfer(transfers[0].transfer_info) - assert manager.current_placement[0][0] == [4, 2] + assert manager.current_placement[0][0] == [0, 1, 2, 4, 2] manager.active_transfer = transfers[1] manager._commit_transfer(transfers[1].transfer_info) @@ -1221,8 +1226,14 @@ def is_finished(self): manager._weights = [object(), object()] manager.world_size = 4 manager.pending_transfer_infos = [remote_info, local_info0, local_info1] - manager.target_placement = [[[2]], [[4]]] - manager.current_placement = [[[0]], [[3]]] + manager.target_placement = [ + [[0, 1, 2], [2, 3, 3], [4, 5, 2], [6, 7, 0]], + [[0, 1, 4], [2, 3, 5], [4, 5, 6], [6, 7, 1]], + ] + manager.current_placement = [ + [[0, 1, 4], [2, 3, 5], [4, 5, 6], [6, 7, 0]], + [[0, 1, 2], [2, 3, 4], [4, 5, 6], [6, 7, 0]], + ] manager.global_rank = 1 manager.state = manager_module.EPLBManagerState.TRANSFERRING committed = [] @@ -1279,7 +1290,10 @@ def all_gather_object(output, local_state, **_kwargs): manager._step_transferring() assert manager.state is manager_module.EPLBManagerState.COLLECTING - assert manager.current_placement == [[[2]], [[4]]] + assert manager.current_placement == [ + [[0, 1, 2], [2, 3, 3], [4, 5, 2], [6, 7, 0]], + [[0, 1, 4], [2, 3, 5], [4, 5, 6], [6, 7, 1]], + ] assert not hasattr(manager, "pending_transfer_infos") assert not hasattr(manager, "target_placement") assert not hasattr(manager, "rebalance_started_at") @@ -1334,24 +1348,24 @@ def is_finished(self): original_overlap_stream = g_infer_context.overlap_stream manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - transfer_info = EPLBTransferInfo(0, 0, 2, 0, 0) + transfer_info = EPLBTransferInfo(0, 0, 0, 0, 0) transfer = Transfer(live, received, transfer_info) manager.active_transfer = transfer manager.active_transfer_batch = [transfer_info] manager.control_group = object() manager.world_size = 1 manager.pending_transfer_infos = [transfer_info] - manager.num_primary_experts_per_rank = 0 + manager.num_primary_experts_per_rank = 1 manager.num_logical_experts = 1 manager.global_rank = 0 - manager.target_placement = [[[2]]] + manager.target_placement = [[[0]]] manager._eplb_impls = [ SimpleNamespace( - local_logics_expert_ids_list=[1], + local_logics_expert_ids_list=[0], logical_to_physical_map=torch.zeros((1, 1), dtype=torch.int32, device="cuda"), ) ] - manager.current_placement = [[[1]]] + manager.current_placement = [[[0]]] manager.rebalance_started_at = time.time() manager.state = manager_module.EPLBManagerState.TRANSFERRING monkeypatch.setattr( @@ -1450,8 +1464,14 @@ def test_manager_enters_transferring_state_with_planned_work(monkeypatch): manager.transfer_group = object() manager.world_size = 2 manager.num_logical_experts = 4 - placement = [[[2], [3]], [[0], [1]]] - manager.current_placement = [[[3], [2]], [[1], [0]]] + placement = [ + [[0, 1, 2], [2, 3, 3]], + [[0, 1, 3], [2, 3, 0]], + ] + manager.current_placement = [ + [[0, 1, 3], [2, 3, 2]], + [[0, 1, 2], [2, 3, 1]], + ] monkeypatch.setattr( manager_module.dist, "broadcast_object_list", @@ -1543,7 +1563,7 @@ def start(self): self.started = True manager.planner = object() - manager.current_placement = [[[1]]] + manager.current_placement = [[[0, 1, 2, 3]]] manager._publish_expert_load_metric = lambda _global_load: None monkeypatch.setattr(manager_module, "EPLBPlanTask", PlanTask) @@ -1594,7 +1614,7 @@ def test_manager_planning_without_changes_returns_to_collecting(monkeypatch): manager.steps = 11 manager.step_interval = 20 manager.next_evaluation_step = 31 - manager.current_placement = [[[1], [0]]] + manager.current_placement = [[[0, 1], [1, 0]]] result = None def broadcast(values, **_kwargs): @@ -1605,7 +1625,7 @@ def broadcast(values, **_kwargs): manager.step() assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH - result = [[[1], [0]]] + result = [[[0, 1], [1, 0]]] manager.step() assert manager.state is manager_module.EPLBManagerState.COLLECTING @@ -1823,9 +1843,9 @@ def new_group(*args, **kwargs): monkeypatch.setattr(manager_module.dist, "new_group", new_group) all_gather_calls = [] - def all_gather_object(output, local_redundant_expert_ids_by_layer, group): - all_gather_calls.append((local_redundant_expert_ids_by_layer, group)) - output[:] = [local_redundant_expert_ids_by_layer, [[0, 1]]] + def all_gather_object(output, local_expert_ids_by_layer, group): + all_gather_calls.append((local_expert_ids_by_layer, group)) + output[:] = [local_expert_ids_by_layer, [[2, 3, 0, 1]]] monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) logs = [] @@ -1836,9 +1856,9 @@ def all_gather_object(output, local_redundant_expert_ids_by_layer, group): assert manager.state is manager_module.EPLBManagerState.COLLECTING assert (manager.control_group, manager.transfer_group) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 2 - assert all_gather_calls == [([[2, 3]], groups[0])] + assert all_gather_calls == [([[0, 1, 2, 3]], groups[0])] assert manager.planner.rebalance_gain_threshold == 0.07 - assert manager.current_placement == [[[2, 3], [0, 1]]] + assert manager.current_placement == [[[0, 1, 2, 3], [2, 3, 0, 1]]] assert manager.metric_client is metric_client assert metric_client_ports == [1234] assert manager.next_evaluation_step == manager.step_interval diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index da05afb515..14c352a1b2 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -75,8 +75,8 @@ def _worker(rank, port): transfer_group = dist.new_group([0, 1], backend="gloo") weights = [_FakeWeight(rank, layer_index) for layer_index in range(2)] - current = [[2], [0]] - target = [[3], [1]] + current = [[0, 1, 2], [2, 3, 0]] + target = [[0, 1, 3], [2, 3, 1]] for expected_layer in range(2): transfer_infos = build_transfer_plan( current, From 901e531cae9ea9ecbf78366f171c86c8061db94a Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 14:43:08 +0000 Subject: [PATCH 40/72] simplify EPLB greedy placement planning --- .../meta_weights/fused_moe/eplb_placement.py | 9 + .../meta_weights/fused_moe/eplb_planner.py | 276 +++++++----------- .../meta_weights/fused_moe/impl/__init__.py | 5 + .../fused_moe/impl/deepgemm_impl.py | 9 +- lightllm/server/metrics/metrics.py | 2 +- .../mode_backend/chunked_prefill/impl.py | 4 +- .../mode_backend/dp_backend/impl.py | 4 +- .../model_infer/mode_backend/eplb_manager.py | 17 +- lightllm/utils/envs_utils.py | 24 +- unit_tests/common/fused_moe/test_eplb.py | 30 +- 10 files changed, 161 insertions(+), 219 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py index d043e4cd6f..1792f0f803 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -1,3 +1,12 @@ +"""构建完整专家布局及其紧凑路由元数据。 + +本模块统一使用 +``[layer][rank][local physical expert] -> logical expert`` 表示专家布局; +处理单层布局的函数会省略 layer 维。每个 rank 的固定主专家排在前面, +可迁移的冗余专家排在后面。 +""" + + def build_initial_local_expert_ids( num_logical_experts: int, num_ranks: int, diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py index 7640a37162..e5ed87fa84 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py @@ -1,13 +1,10 @@ -"""Pure-Python redundant-expert placement planning for EPLB. +"""使用纯 Python 实现 EPLB 冗余专家布局规划。 -The planner deliberately uses nested lists instead of tensors. Tensor -conversion belongs to the manager's distributed-communication and migration -boundaries; keeping it out of this module makes placement algorithms easy to -read, test, and replace. +规划器有意使用嵌套 list,而不是 Tensor。Tensor 转换仅发生在 manager 的 +分布式通信和迁移边界;规划模块不依赖 Tensor,更易于阅读、测试和替换算法。 """ from abc import ABC, abstractmethod -from collections import Counter from math import ceil from typing import List @@ -21,7 +18,7 @@ class EPLBPlanner(ABC): - """Interface for planning redundant expert placement.""" + """冗余专家布局规划接口。""" @abstractmethod def plan( @@ -29,11 +26,17 @@ def plan( logical_expert_load: LogicalExpertLoad, current_placement: ExpertPlacement, ) -> ExpertPlacement: - """Return a concrete ``[layer][rank][local physical expert]`` placement.""" + """返回完整的 ``[layer][rank][local physical expert]`` 专家布局。""" class GreedyEPLBPlanner(EPLBPlanner): - """Greedy planner based on global logical-expert loads.""" + """逐个填充冗余槽位,尽量降低最繁忙 rank 的负载。 + + 主专家始终固定不动。每轮枚举所有合法的 ``(rank, logical_expert)`` 组合, + 选择加入副本后最大 rank 负载最小的候选,直至填满全部冗余槽位。 + 新增副本只会改变对应逻辑专家的负载分摊,因此可以直接根据当前布局评估 + 每个候选,不再需要单独分配副本数量或执行回溯搜索。 + """ def __init__( self, @@ -41,7 +44,6 @@ def __init__( num_redundant_experts_per_rank: int, *, expert_alignment: int = 1, - rebalance_gain_threshold: float = 0.0, ): if world_size <= 1: raise ValueError("world_size must be greater than one") @@ -49,25 +51,21 @@ def __init__( raise ValueError("num_redundant_experts_per_rank must be positive") if expert_alignment <= 0: raise ValueError("expert_alignment must be positive") - if not 0.0 <= rebalance_gain_threshold <= 1.0: - raise ValueError("rebalance_gain_threshold must be between 0.0 and 1.0") self.world_size = world_size self.num_redundant_experts_per_rank = num_redundant_experts_per_rank self.expert_alignment = expert_alignment - self.rebalance_gain_threshold = rebalance_gain_threshold def plan( self, logical_expert_load: LogicalExpertLoad, current_placement: ExpertPlacement, ) -> ExpertPlacement: - """Return a new concrete placement. + """根据全局逻辑专家负载生成新的完整布局。 - ``logical_expert_load`` is ``[layer][logical_expert]`` and contains - the load summed across all ranks. - ``current_placement`` is ``[layer][rank][local_physical_expert]``. - Each rank row contains its fixed primary experts followed by its - movable redundant experts. The returned placement has the same shape. + ``logical_expert_load`` 的形状为 ``[layer][logical_expert]``,保存所有 + rank 汇总后的负载。``current_placement`` 的形状为 + ``[layer][rank][local_physical_expert]``;每个 rank 的固定主专家在前, + 可迁移冗余专家在后。返回布局与当前布局形状相同。 """ load = [[float(value) for value in layer] for layer in logical_expert_load] current = [[[int(expert) for expert in rank] for rank in layer] for layer in current_placement] @@ -78,11 +76,12 @@ def plan( candidate_rank_load = self.estimate_rank_load(load, candidates) placement = [] for layer, candidate in enumerate(candidates): + # 每层独立决定是否采用候选布局。只有最大 rank 负载严格下降时 + # 才迁移,避免相同负载下的无收益调整,也使零负载层保持原布局。 before = max(before_rank_load[layer]) after = max(candidate_rank_load[layer]) - gain = (before - after) / max(before, 1.0) - changed = candidate != current[layer] and gain > self.rebalance_gain_threshold - placement.append(candidate if changed else current[layer]) + improved = candidate != current[layer] and after < before + placement.append(candidate if improved else current[layer]) return placement def estimate_rank_load( @@ -90,7 +89,7 @@ def estimate_rank_load( logical_expert_load: LogicalExpertLoad, placement: ExpertPlacement, ) -> RankLoad: - """Estimate aligned physical-expert work for each layer and rank.""" + """估算每层、每个 rank 上经过对齐后的物理专家工作量。""" load = [[float(value) for value in layer] for layer in logical_expert_load] normalized_placement = [[[int(expert) for expert in rank] for rank in layer] for layer in placement] num_logical_experts = self._validate_inputs(load, normalized_placement) @@ -99,7 +98,7 @@ def estimate_rank_load( locations = self._expert_locations(layer_placement, num_logical_experts) rank_load = [0.0] * self.world_size for expert, expert_locations in enumerate(locations): - for rank, value in enumerate(self._physical_load(layer_load[expert], expert_locations)): + for rank, value in enumerate(self._rank_load_for_expert(layer_load[expert], expert_locations)): rank_load[rank] += value estimated.append(rank_load) return estimated @@ -110,175 +109,96 @@ def _plan_layer( current_placement: List[List[int]], ) -> List[List[int]]: num_logical_experts = len(logical_load) - experts_per_rank = num_logical_experts // self.world_size - owner = [expert // experts_per_rank for expert in range(num_logical_experts)] - copy_count = self._allocate_copy_count(logical_load, owner) - current_redundant_placement = [row[experts_per_rank:] for row in current_placement] + num_primary_experts_per_rank = num_logical_experts // self.world_size + current_redundant_placement = [row[num_primary_experts_per_rank:] for row in current_placement] - redundant_placement = [[-1] * self.num_redundant_experts_per_rank for _ in range(self.world_size)] - locations = [{owner_rank} for owner_rank in owner] - expert_rank_load = [ - self._physical_load( - logical_load[expert], - locations[expert], - ) - for expert in range(num_logical_experts) + # 从不可变的主专家布局开始。只要目标 rank 尚未持有该专家副本, + # 对应的 (rank, expert) 组合就是合法候选。 + redundant_placement: List[List[int]] = [[] for _ in range(self.world_size)] + locations = [{expert // num_primary_experts_per_rank} for expert in range(num_logical_experts)] + rank_load_by_expert = [ + self._rank_load_for_expert(logical_load[expert], locations[expert]) for expert in range(num_logical_experts) ] rank_load = [ - sum(expert_rank_load[expert][rank] for expert in range(num_logical_experts)) + sum(rank_load_by_expert[expert][rank] for expert in range(num_logical_experts)) for rank in range(self.world_size) ] - instances = sorted( - [expert for expert, copies in enumerate(copy_count) for _ in range(copies - 1)], - key=lambda expert: ( - -copy_count[expert], - -logical_load[expert] / copy_count[expert], - expert, - ), - ) - for index, expert in enumerate(instances): - best = None + total_redundant_experts = self.world_size * self.num_redundant_experts_per_rank + for _ in range(total_redundant_experts): + best_score = None + best_rank = None + best_expert = None + best_rank_load = None + best_expert_rank_load = None for rank in range(self.world_size): - if rank in locations[expert]: - continue - empty_slots = [slot for slot, value in enumerate(redundant_placement[rank]) if value < 0] - if not empty_slots: - continue - slot = min( - empty_slots, - key=lambda candidate: ( - current_redundant_placement[rank][candidate] != expert, - candidate, - ), - ) - trial_redundant_placement = [row[:] for row in redundant_placement] - trial_redundant_placement[rank][slot] = expert - trial_locations = [set(ranks) for ranks in locations] - trial_locations[expert].add(rank) - if not self._can_complete( - instances[index + 1 :], - trial_redundant_placement, - trial_locations, - owner, - ): + if len(redundant_placement[rank]) == self.num_redundant_experts_per_rank: continue - next_expert_load = self._physical_load( - logical_load[expert], - trial_locations[expert], - ) - trial_rank_load = [ - value - expert_rank_load[expert][target_rank] + next_expert_load[target_rank] - for target_rank, value in enumerate(rank_load) - ] - candidate = ( - max(trial_rank_load), - current_redundant_placement[rank][slot] != expert, - sum(trial_rank_load), - rank, - slot, - trial_rank_load, - next_expert_load, - ) - if best is None or candidate[:5] < best[:5]: - best = candidate + for expert in range(num_logical_experts): + if rank in locations[expert]: + continue - if best is None: - raise RuntimeError("EPLB planner found no valid redundant expert placement") - _, _, _, rank, slot, rank_load, next_expert_load = best - redundant_placement[rank][slot] = expert - locations[expert].add(rank) - expert_rank_load[expert] = next_expert_load + next_locations = locations[expert] | {rank} + next_expert_rank_load = self._rank_load_for_expert(logical_load[expert], next_locations) + next_rank_load = [ + load - rank_load_by_expert[expert][target_rank] + next_expert_rank_load[target_rank] + for target_rank, load in enumerate(rank_load) + ] - return [ - list(range(rank * experts_per_rank, (rank + 1) * experts_per_rank)) + redundant_experts - for rank, redundant_experts in enumerate(redundant_placement) - ] - - def _allocate_copy_count(self, logical_load: List[float], owner: List[int]) -> List[int]: - copy_count = [1] * len(logical_load) - replicas_by_owner = [0] * self.world_size - owner_capacity = self.num_redundant_experts_per_rank * (self.world_size - 1) - total_replicas = self.num_redundant_experts_per_rank * self.world_size - for _ in range(total_replicas): - candidates = [ - expert - for expert in range(len(logical_load)) - if copy_count[expert] < self.world_size and replicas_by_owner[owner[expert]] < owner_capacity - ] - if not candidates: - raise RuntimeError("EPLB planner cannot allocate all redundant copies") - expert = max( - candidates, - key=lambda candidate: ( - logical_load[candidate] / copy_count[candidate], - -candidate, - ), - ) - copy_count[expert] += 1 - replicas_by_owner[owner[expert]] += 1 - return copy_count + # 首先最小化最大 rank 负载;负载相同时依次选择总对齐工作量 + # 更小、能保留现有本地副本的候选。最后使用 rank/expert ID + # 打破平局,保证相同输入始终得到相同结果。 + score = ( + max(next_rank_load), + sum(next_rank_load), + expert not in current_redundant_placement[rank], + rank, + expert, + ) + if best_score is None or score < best_score: + best_score = score + best_rank = rank + best_expert = expert + best_rank_load = next_rank_load + best_expert_rank_load = next_expert_rank_load - def _can_complete( - self, - remaining_instances: List[int], - redundant_placement: List[List[int]], - locations: List[set], - owner: List[int], - ) -> bool: - """Check that a greedy choice leaves a legal assignment for all slots.""" - remaining = Counter(remaining_instances) - capacity = [sum(expert < 0 for expert in row) for row in redundant_placement] - memo = set() + # 输入校验已保证每个 rank 都有足够多且互不重复的非本地主专家, + # 因此所有冗余槽位一定可以填满。 + assert best_rank is not None and best_expert is not None + assert best_rank_load is not None and best_expert_rank_load is not None + redundant_placement[best_rank].append(best_expert) + locations[best_expert].add(best_rank) + rank_load = best_rank_load + rank_load_by_expert[best_expert] = best_expert_rank_load - def search() -> bool: - if not remaining: - return True - state = ( - tuple(sorted(remaining.items())), - tuple(capacity), - tuple(tuple(sorted(ranks)) for ranks in locations), - ) - if state in memo: - return False - memo.add(state) - - expert = min( - remaining, - key=lambda item: ( - sum(capacity[rank] > 0 and rank not in locations[item] for rank in range(self.world_size)) - - remaining[item], - -remaining[item], - item, - ), - ) - candidate_ranks = [ - rank - for rank in range(self.world_size) - if capacity[rank] > 0 and rank not in locations[expert] and rank != owner[expert] + # 布局质量只取决于选中了哪些专家,与它们在本地冗余槽中的顺序无关。 + # 已经选中的现有专家继续使用原槽位,只有新增专家才填入剩余槽位; + # 这样无需干扰上面的负载均衡循环,也能尽量减少权重传输。 + for rank, selected_experts in enumerate(redundant_placement): + selected_set = set(selected_experts) + new_experts = iter(expert for expert in selected_experts if expert not in current_redundant_placement[rank]) + redundant_placement[rank] = [ + expert if expert in selected_set else next(new_experts) for expert in current_redundant_placement[rank] ] - if len(candidate_ranks) < remaining[expert]: - return False - count = remaining.pop(expert) - if count > 1: - remaining[expert] = count - 1 - for rank in sorted(candidate_ranks, key=lambda item: (-capacity[item], item)): - capacity[rank] -= 1 - locations[expert].add(rank) - if search(): - locations[expert].remove(rank) - capacity[rank] += 1 - remaining[expert] = count - return True - locations[expert].remove(rank) - capacity[rank] += 1 - remaining[expert] = count - return False - return search() + return [ + list( + range( + rank * num_primary_experts_per_rank, + (rank + 1) * num_primary_experts_per_rank, + ) + ) + + redundant_experts + for rank, redundant_experts in enumerate(redundant_placement) + ] - def _physical_load(self, logical_expert_load: float, locations: set) -> List[float]: + def _rank_load_for_expert( + self, + logical_expert_load: float, + locations: set, + ) -> List[float]: + """返回该专家在各 rank 上经过对齐后的负载贡献。""" physical_expert_load = logical_expert_load / len(locations) aligned_load = ceil(physical_expert_load / self.expert_alignment) * self.expert_alignment result = [0.0] * self.world_size @@ -291,6 +211,7 @@ def _expert_locations( placement: List[List[int]], num_logical_experts: int, ) -> List[set]: + """反转单层布局,并拒绝同一 rank 上的重复专家副本。""" locations = [set() for _ in range(num_logical_experts)] for rank, row in enumerate(placement): for expert in row: @@ -304,6 +225,7 @@ def _validate_inputs( logical_expert_load: LogicalExpertLoad, placement: ExpertPlacement, ) -> int: + """校验完整布局约束,并返回逻辑专家数量。""" if not logical_expert_load: raise ValueError("logical_expert_load must contain at least one layer") if len(placement) != len(logical_expert_load): diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 20522b5c37..44c3091903 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -12,6 +12,11 @@ def create_fuse_moe_impl( quant_method: QuantizationMethod, enable_ep_moe: bool = False, ): + """创建持有自身路由运行态的 MoE 执行实现。 + + 这里直接返回完成初始化的对象,而不是仅返回实现类,使 EPLB 布局、路由 + 计数器等后端专属状态与使用它们的 kernel 保持在同一个实现对象中。 + """ if enable_ep_moe: impl_cls = FuseMoeDeepGEMM elif quant_method.method_name == "awq_marlin": 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 e19bb63634..1122d06058 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 @@ -36,6 +36,12 @@ def __init__(self, *args, **kwargs): self._init_eplb_runtime() def _init_eplb_runtime(self): + """初始化本地物理槽位以及可更新的 EPLB 路由运行态。 + + ``local_logics_expert_ids_list`` 始终描述全部本地物理行:固定主专家行 + 在前,冗余专家行在后。负载均衡只允许替换冗余后缀,并在同一个安全 + 推理边界同时更新专家权重行和 ``logical_to_physical_map``。 + """ world_size = get_global_world_size() assert self.n_routed_experts % world_size == 0 global_rank = get_global_rank() @@ -58,6 +64,7 @@ def _init_eplb_runtime(self): ), dtype=torch.int32, ).cuda() + # 始终按逻辑专家统计负载,冗余副本不会拆散规划器观察到的负载信号。 self.route_counter = torch.zeros(self.n_routed_experts, dtype=torch.int64, device="cuda") self.recording = True else: @@ -84,7 +91,7 @@ def _select_experts( scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, ): - """Select logical experts without applying the EPLB physical layout.""" + """只选择逻辑专家,不在此阶段应用 EPLB 物理布局。""" from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts topk_weights, topk_ids = select_experts( diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index 9b8c620322..0d53669a16 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -34,7 +34,7 @@ "lightllm_num_running_reqs": "Number of running requests", "lightllm_eplb_topk_expert_imbalance_ratio": ( "Maximum routed token count divided by the mean across logical experts, averaged across MoE layers in the " - "latest EPLB sample window" + "accumulated EPLB routing sample" ), } diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 28ef888b4e..50fc80cd54 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -63,7 +63,9 @@ def infer_loop(self): self._try_read_new_reqs() - # Keep EPLB collectives ordered before normal collectives/forward. + # EPLB step 可能发起所有 rank 都必须按相同顺序参与的控制面 + # collective。固定放在请求读取之后、常规通信和 forward 之前, + # 即使本 rank 当前没有请求,也不会与后续 collective 交错。 if self.eplb_manager is not None: self.eplb_manager.step() diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 387da3d07f..ace0d2b617 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -124,7 +124,9 @@ def infer_loop(self): self._try_read_new_reqs() - # Keep EPLB collectives ordered before normal collectives/forward. + # EPLB step 可能发起所有 rank 都必须按相同顺序参与的控制面 + # collective。固定放在请求读取之后、常规通信和 forward 之前, + # 即使本 rank 当前没有请求,也不会与后续 collective 交错。 if self.eplb_manager is not None: self.eplb_manager.step() diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 5d584ff5fa..f837c7d51e 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -28,10 +28,7 @@ get_global_rank, get_global_world_size, ) -from lightllm.utils.envs_utils import ( - get_eplb_rebalance_gain_threshold, - get_eplb_step_interval, -) +from lightllm.utils.envs_utils import get_eplb_step_interval from lightllm.utils.log_utils import init_logger from lightllm.utils.shm_port_args import get_shm_port_args @@ -58,7 +55,7 @@ class EPLBManager: ``COLLECTING -> EVALUATING -> PLANNING -> WAIT_PLAN_FINISH -> TRANSFERRING -> COLLECTING`` - 当收集的平均专家 token 数不足时,``EVALUATING`` 会回到 + 当累计的平均专家 token 数不足时,``EVALUATING`` 会回到 ``COLLECTING``;当规划器认为无需调整布局时,``WAIT_PLAN_FINISH`` 会 回到 ``COLLECTING``。每次调用 :meth:`step` 最多推进一个状态,布局 规划和权重传输在后台执行,主推理线程负责评估、轮询和提交结果。 @@ -80,7 +77,8 @@ def __init__(self, model: TpPartBaseModel) -> None: self.num_redundant_experts_per_rank: int = first_impl.num_redundant_experts_per_rank self.num_primary_experts_per_rank: int = self.num_logical_experts // self.world_size - # 评估调度:steps 只在 COLLECTING 状态递增。 + # 评估调度:steps 只在 COLLECTING 状态递增。route counter 不按周期 + # 清零,让低流量服务可以跨多个评估周期积累到足够可靠的样本量。 self.step_interval: int = get_eplb_step_interval() self.steps: int = 0 @@ -109,7 +107,6 @@ def __init__(self, model: TpPartBaseModel) -> None: self.world_size, self.num_redundant_experts_per_rank, expert_alignment=EPLB_EXPERT_ALIGNMENT, - rebalance_gain_threshold=get_eplb_rebalance_gain_threshold(), ) self.state = EPLBManagerState.COLLECTING @@ -150,7 +147,7 @@ def step(self) -> None: # 状态处理:与 step() 的分发顺序保持一致。 def _step_collecting(self) -> None: - """记录一个采样步,并在采样窗口结束后进入评估状态。""" + """记录一个采样步,并在当前评估周期结束后进入评估状态。""" self.steps += 1 if self.steps < self.next_evaluation_step: return @@ -165,6 +162,8 @@ def _step_evaluating(self) -> None: raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") # 将各层累计的路由计数复制到 CPU,后续规划统一使用这份快照。 + # 此处有意不清零 GPU counter:下一周期继续累计,而本轮异步规划 + # 使用独立的 CPU 快照,不会与推理线程后续的 atomic add 竞争。 local_load = torch.stack([counter.detach().cpu() for counter in counters]) # 汇集各 rank 的 token 总数,判断当前统计量是否足以进行布局规划。 @@ -411,7 +410,7 @@ def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: - """Average each layer's maximum-to-mean logical-expert token ratio.""" + """计算各层逻辑专家最大 token 数与平均值之比,再对所有层取平均。""" if global_load.ndim != 2: raise ValueError("global_load must be [layers, logical_experts]") global_load = global_load.to(torch.float64) diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index a647e344b6..2521fa8b36 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -98,35 +98,13 @@ def get_lightllm_websocket_max_message_size(): @lru_cache(maxsize=None) def get_eplb_step_interval(): - """Return the number of inference steps between EPLB attempts.""" + """返回两次 EPLB 评估之间的推理步数。""" interval = int(os.getenv("LIGHTLLM_EPLB_STEP_INTERVAL", 20)) if interval <= 0: raise ValueError("LIGHTLLM_EPLB_STEP_INTERVAL must be greater than 0") return interval -@lru_cache(maxsize=None) -def get_eplb_rebalance_gain_threshold() -> float: - """Return the EPLB gain threshold: estimated critical-load reduction ratio; 0.05 means 5%.""" - env_name = "LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD" - raw_value = os.getenv(env_name, "0.05") - value = float(raw_value) - if not 0.0 <= value <= 1.0: - raise ValueError(f"{env_name} must be a ratio between 0.0 and 1.0, got {raw_value!r}") - return value - - -@lru_cache(maxsize=None) -def get_eplb_placement_stickiness() -> float: - """Return the EPLB placement stickiness: a keep-bonus, as a fraction of the mean per-layer expert load.""" - env_name = "LIGHTLLM_EPLB_PLACEMENT_STICKINESS" - raw_value = os.getenv(env_name, "0.1") - value = float(raw_value) - if not 0.0 <= value <= 1.0: - raise ValueError(f"{env_name} must be a ratio between 0.0 and 1.0, got {raw_value!r}") - return value - - @lru_cache(maxsize=None) def get_triton_autotune_level(): return int(os.getenv("LIGHTLLM_TRITON_AUTOTUNE_LEVEL", 0)) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index bd3073fdc8..0dd3b33ad5 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -298,7 +298,6 @@ def test_eplb_planner_builds_legal_concrete_slot_layout(): 4, 1, expert_alignment=1, - rebalance_gain_threshold=0.0, ) current = _initial_expert_placement(8, 4, 1).unsqueeze(0).tolist() load = torch.ones((1, 4, 8), dtype=torch.int64) @@ -342,11 +341,32 @@ def test_eplb_planner_does_not_move_zero_load_experts(): assert result == current -def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): +def test_eplb_planner_iteratively_places_hot_expert_on_idle_rank(): + planner = GreedyEPLBPlanner(2, 1) + current = [[[0, 1, 3], [2, 3, 1]]] + + result = planner.plan([[1000, 1, 1, 1]], current) + + assert result == [[[0, 1, 3], [2, 3, 0]]] + assert planner.estimate_rank_load([[1000, 1, 1, 1]], result) == [[502.0, 502.0]] + + +def test_eplb_planner_keeps_selected_experts_in_their_current_slots(): + planner = GreedyEPLBPlanner(4, 2) + current = _initial_expert_placement(8, 4, 2).unsqueeze(0).tolist() + + result = planner.plan([[50, 98, 54, 6, 34, 66, 63, 52]], current) + + # Rank 0 的专家 3 保留在原来的第二个冗余槽位,仅将第一个槽位 + # 从专家 2 替换为专家 4。 + assert current[0][0][2:] == [2, 3] + assert result[0][0][2:] == [4, 3] + + +def test_eplb_planner_fills_every_rank_with_distinct_nonlocal_experts(): planner = GreedyEPLBPlanner( 4, 1, - rebalance_gain_threshold=0.0, ) current = _initial_expert_placement(16, 4, 1).unsqueeze(0).tolist() load = torch.randint( @@ -365,7 +385,7 @@ def test_eplb_planner_reserves_rank_capacity_for_remaining_copies(): assert all(expert // 4 != rank for expert in row[4:]) -def test_eplb_planner_keeps_high_redundancy_search_state_isolated(): +def test_eplb_planner_supports_multiple_redundant_experts_per_rank(): planner = GreedyEPLBPlanner(4, 3) current = _initial_expert_placement(16, 4, 3).unsqueeze(0).tolist() load = [ @@ -1826,7 +1846,6 @@ def test_manager_initializes_without_transfer_task(monkeypatch): monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) monkeypatch.setattr(manager_module, "get_eplb_step_interval", lambda: 20) - monkeypatch.setattr(manager_module, "get_eplb_rebalance_gain_threshold", lambda: 0.07) monkeypatch.setattr(manager_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=1234)) metric_client = SimpleNamespace() metric_client_ports = [] @@ -1857,7 +1876,6 @@ def all_gather_object(output, local_expert_ids_by_layer, group): assert (manager.control_group, manager.transfer_group) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 2 assert all_gather_calls == [([[0, 1, 2, 3]], groups[0])] - assert manager.planner.rebalance_gain_threshold == 0.07 assert manager.current_placement == [[[0, 1, 2, 3], [2, 3, 0, 1]]] assert manager.metric_client is metric_client assert metric_client_ports == [1234] From 7a8fee400bb47549cc7cbbc78e97b0ac141f23cb Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 18 Sep 2026 14:53:53 +0000 Subject: [PATCH 41/72] reset EPLB route counters across placements --- .../model_infer/mode_backend/eplb_manager.py | 22 +++++++++++--- unit_tests/common/fused_moe/test_eplb.py | 29 +++++++++++++++++++ 2 files changed, 47 insertions(+), 4 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index f837c7d51e..b90df92e70 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -77,8 +77,8 @@ def __init__(self, model: TpPartBaseModel) -> None: self.num_redundant_experts_per_rank: int = first_impl.num_redundant_experts_per_rank self.num_primary_experts_per_rank: int = self.num_logical_experts // self.world_size - # 评估调度:steps 只在 COLLECTING 状态递增。route counter 不按周期 - # 清零,让低流量服务可以跨多个评估周期积累到足够可靠的样本量。 + # 评估调度:steps 只在 COLLECTING 状态递增。route counter 从当前 + # 布局生效时开始累计,让低流量服务可以跨多个评估周期收集足够样本。 self.step_interval: int = get_eplb_step_interval() self.steps: int = 0 @@ -111,6 +111,7 @@ def __init__(self, model: TpPartBaseModel) -> None: self.state = EPLBManagerState.COLLECTING self.next_evaluation_step = self.step_interval + self._clear_route_counters() if self.global_rank == 0: self.metric_client: MetricClient = MetricClient(get_shm_port_args().metric_port) @@ -162,8 +163,9 @@ def _step_evaluating(self) -> None: raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") # 将各层累计的路由计数复制到 CPU,后续规划统一使用这份快照。 - # 此处有意不清零 GPU counter:下一周期继续累计,而本轮异步规划 - # 使用独立的 CPU 快照,不会与推理线程后续的 atomic add 竞争。 + # 此处有意不清零 GPU counter:如果样本不足或无需迁移,下一周期会 + # 继续累计;成功切换到新布局后才重新开始统计。本轮异步规划使用独立 + # 的 CPU 快照,不会与推理线程后续的 atomic add 竞争。 local_load = torch.stack([counter.detach().cpu() for counter in counters]) # 汇集各 rank 的 token 总数,判断当前统计量是否足以进行布局规划。 @@ -278,6 +280,7 @@ def _step_transferring(self) -> None: if not transfer_batch: self.current_placement = self.target_placement elapsed = time.time() - self.rebalance_started_at + self._clear_route_counters() del self.pending_transfer_infos del self.target_placement del self.rebalance_started_at @@ -399,6 +402,17 @@ def _publish_expert_load_metric(self, global_load: torch.Tensor) -> None: _expert_load_imbalance_ratio(global_load), ) + def _clear_route_counters(self) -> None: + """在 overlap stream 上清空所有层的逻辑专家路由计数。""" + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + # route counter 由 forward 中的 Triton kernel 在 overlap stream 上更新。 + # 将 zero_ 排到同一条 stream,可保证它位于此前 forward 之后、下一次 + # forward 之前,无需额外 synchronize,也不会与 atomic add 并发。 + with torch.cuda.stream(g_infer_context.get_overlap_stream()): + for impl in self._eplb_impls: + impl.route_counter.zero_() + def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: weights_by_id: Dict[int, FusedMoeWeight] = {} diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 0dd3b33ad5..06f7722744 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1257,7 +1257,9 @@ def is_finished(self): manager.global_rank = 1 manager.state = manager_module.EPLBManagerState.TRANSFERRING committed = [] + cleared_route_counters = [] manager._commit_transfer = committed.append + manager._clear_route_counters = lambda: cleared_route_counters.append(True) waits = [] overlap_stream = object() monkeypatch.setattr( @@ -1317,6 +1319,7 @@ def all_gather_object(output, local_state, **_kwargs): assert not hasattr(manager, "pending_transfer_infos") assert not hasattr(manager, "target_placement") assert not hasattr(manager, "rebalance_started_at") + assert cleared_route_counters == [True] assert local_states == [ state(local_info0, False), state(local_info0, True), @@ -1823,6 +1826,25 @@ def test_manager_requires_more_than_one_rank(monkeypatch): manager_module.EPLBManager(type("Model", (), {})()) +def test_manager_clears_all_route_counters_on_overlap_stream(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + counters = [torch.tensor([1, 2]), torch.tensor([3, 4])] + manager._eplb_impls = [SimpleNamespace(route_counter=counter) for counter in counters] + overlap_stream = object() + used_streams = [] + monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) + monkeypatch.setattr( + manager_module.torch.cuda, + "stream", + lambda stream: (used_streams.append(stream) or nullcontext()), + ) + + manager._clear_route_counters() + + assert used_streams == [overlap_stream] + assert all(torch.count_nonzero(counter) == 0 for counter in counters) + + def test_manager_initializes_without_transfer_task(monkeypatch): weight = type( "Weight", @@ -1846,6 +1868,12 @@ def test_manager_initializes_without_transfer_task(monkeypatch): monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) monkeypatch.setattr(manager_module, "get_eplb_step_interval", lambda: 20) + clear_calls = [] + monkeypatch.setattr( + manager_module.EPLBManager, + "_clear_route_counters", + lambda manager: clear_calls.append(manager), + ) monkeypatch.setattr(manager_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=1234)) metric_client = SimpleNamespace() metric_client_ports = [] @@ -1880,6 +1908,7 @@ def all_gather_object(output, local_expert_ids_by_layer, group): assert manager.metric_client is metric_client assert metric_client_ports == [1234] assert manager.next_evaluation_step == manager.step_interval + assert clear_calls == [manager] assert "planner=GreedyEPLBPlanner" in logs[0] assert weight.fuse_moe_impl.recording assert manager._eplb_impls[0] is weight.fuse_moe_impl From d89024d6b24d6b766018a2db67c6b4e0a570830c Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 20 Sep 2026 07:58:28 +0000 Subject: [PATCH 42/72] refine EPLB placement and transfer planning --- .../meta_weights/fused_moe/eplb_placement.py | 13 +- .../meta_weights/fused_moe/eplb_planner.py | 473 +++++++++++------- .../fused_moe/impl/deepgemm_impl.py | 6 +- .../model_infer/mode_backend/eplb_manager.py | 109 ++-- .../model_infer/mode_backend/eplb_transfer.py | 210 ++++++-- unit_tests/common/fused_moe/test_eplb.py | 304 ++++++++--- .../fused_moe/test_eplb_transfer_gpu.py | 62 ++- 7 files changed, 805 insertions(+), 372 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py index 1792f0f803..dbe098d5df 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -2,8 +2,8 @@ 本模块统一使用 ``[layer][rank][local physical expert] -> logical expert`` 表示专家布局; -处理单层布局的函数会省略 layer 维。每个 rank 的固定主专家排在前面, -可迁移的冗余专家排在后面。 +处理单层布局的函数会省略 layer 维。初始布局按主专家和冗余专家构建; +EPLB 开始运行后,全部物理槽位都可以重新分配。 """ @@ -52,7 +52,7 @@ def build_logical_to_physical_map( ``rank_to_logic_expert_ids`` 的 shape 为 ``[num_ranks, num_physical_experts_per_rank]``,每行包含该 rank 的全部 - 主专家和冗余专家。 + 物理专家。 返回值的 shape 为 ``[num_logical_experts, 2 + routing_slots]``。每一行 对应一个 logical expert:第 0 项是有效副本数,第 1 项 @@ -72,9 +72,8 @@ def build_logical_to_physical_map( num_primary_experts_per_rank = num_logical_experts // num_ranks num_redundant_experts_per_rank = num_physical_experts_per_rank - num_primary_experts_per_rank assert num_redundant_experts_per_rank >= 0 - # 阶段 2:计算固定路由槽宽度。最坏情况下,所有 rank 的全部冗余槽都 - # 指向同一个 logical expert;再加上该 expert 固有的一个主副本,就是 - # 任意 logical expert 可能拥有的最大物理副本数。 + # 阶段 2:计算固定路由槽宽度。该宽度沿用初始化时“一个基础副本加上 + # 全部冗余槽”的容量上界;动态布局不再要求基础副本位于固定槽位。 num_routing_slots = 1 + num_ranks * num_redundant_experts_per_rank assert 0 <= current_rank < num_ranks @@ -165,7 +164,7 @@ def _build_routing_row( 不再依赖 ``current_rank``。列表长度就是该 logical expert 的有效物理 副本数,无需额外传入容易失配的副本数量。 """ - # 阶段 1:候选列表包含一个主副本及全部冗余副本,其长度就是有效副本数。 + # 阶段 1:候选列表包含该专家的全部物理副本,其长度就是有效副本数。 num_valid_replicas = len(physical_expert_ids) assert 0 < num_valid_replicas <= num_routing_slots diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py index e5ed87fa84..44cbd86c9d 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py @@ -1,24 +1,25 @@ -"""使用纯 Python 实现 EPLB 冗余专家布局规划。 +"""使用纯 Python 实现 EPLB 专家布局规划。 规划器有意使用嵌套 list,而不是 Tensor。Tensor 转换仅发生在 manager 的 分布式通信和迁移边界;规划模块不依赖 Tensor,更易于阅读、测试和替换算法。 """ +import heapq from abc import ABC, abstractmethod from math import ceil -from typing import List +from typing import List, Tuple # [layer][logical expert] LogicalExpertLoad = List[List[float]] # [layer][rank][local physical expert] -> logical expert ExpertPlacement = List[List[List[int]]] -# [layer][rank] -RankLoad = List[List[float]] +# (logical expert, replica count, aligned load per replica) +ExpertReplicaGroup = Tuple[int, int, float] class EPLBPlanner(ABC): - """冗余专家布局规划接口。""" + """专家布局规划接口。""" @abstractmethod def plan( @@ -30,12 +31,117 @@ def plan( class GreedyEPLBPlanner(EPLBPlanner): - """逐个填充冗余槽位,尽量降低最繁忙 rank 的负载。 + """使用两阶段启发式算法生成完整的专家布局。 - 主专家始终固定不动。每轮枚举所有合法的 ``(rank, logical_expert)`` 组合, - 选择加入副本后最大 rank 负载最小的候选,直至填满全部冗余槽位。 - 新增副本只会改变对应逻辑专家的负载分摊,因此可以直接根据当前布局评估 - 每个候选,不再需要单独分配副本数量或执行回溯搜索。 + 设计目标 + -------- + 对每一层的全局逻辑专家负载进行快速近似均衡,同时满足以下约束: + + * 每个 rank 的物理专家槽位数量必须固定; + * 每个逻辑专家至少有一个副本; + * 同一个逻辑专家不能在同一个 rank 上出现两次; + * 尽量减少最繁忙 rank 的负载,并保持结果确定,方便规划和迁移测试。 + + ``num_redundant_experts_per_rank`` 个槽位用于复制专家。这里的“冗余专家” + 是布局中的副本槽位,不代表某些主专家槽位不可移动;当前实现允许所有 + 物理槽位重新排列。 + + 单层规划流程 + ------------ + 1. 选择全卡冗余专家:取负载最高的 ``R`` 个逻辑专家,并将它们各放一份 + 到所有 rank。它们在每个 rank 上的负载贡献完全相同,因此在后续比较 + rank 之间的相对负载时可以暂时忽略。 + 2. 决定额外副本数:全卡冗余专家占用了 ``R * world_size`` 个槽位,剩余 + 槽位总数正好是逻辑专家数 ``E``。非冗余专家先各保留一个副本,再将 + 多出的 ``R`` 个副本逐次加给当前 ``load / replica_count`` 最高的专家。 + 这样拆分后的每个副本负载尽量接近。 + 3. 平铺多副本专家:副本数大于 1 的专家按顺序使用循环 rank 游标放置。 + 一组副本数最多为 ``world_size``,所以它会落在互不相同的 rank;连续 + 使用游标还使第一阶段各 rank 的槽位数最多相差一个。 + 4. 放置单副本专家:根据第三步已经形成的 ``rank_load`` 建立最小堆,按 + 单副本负载从高到低取专家,每次分给当前负载最低且仍有空槽的 rank。 + 此阶段只有一个副本的专家,不存在同专家重复约束;rank 填满后从堆中 + 移除。这里负载是第一优先级,剩余槽位数量只用于判断 rank 是否已满。 + 5. 复用当前布局:按当前 rank 顺序执行贪心匹配,每次从尚未使用的候选 + rank 中选择共同专家数量最多的一行,先减少跨 rank 的专家迁移;随后 + 让共同专家继续占用原物理槽位,再减少 rank 内的权重搬运。候选 rank + 重排和本地槽位复用都不会改变规划负载。 + + 示例 + ---- + 假设 ``E=8``、``world_size=4``、``R=1``、``expert_alignment=1``,逻辑 + 专家负载为 ``[40, 12, 9, 8, 7, 5, 4, 3]``。总物理槽位数为 + ``E + R * world_size = 12``,所以每个 rank 必须恰好放置三个专家。 + + 1. 选择全卡冗余专家 + + 专家 0 的负载 40 最高,因此每个 rank 都先放置专家 0。四个副本各 + 分担 ``40 / 4 = 10`` 的负载: + + ``placement = [[0], [0], [0], [0]]`` + + 此时每个 rank 的公共负载都是 10、剩余槽位都是 2。公共负载不会影响 + rank 之间的大小关系,所以后续平衡过程只记录非冗余专家的负载。 + + 2. 计算非冗余专家的副本数 + + 去掉专家 0 后,专家 1 到 7 先各保留一个副本,只能占用 7 个槽位;但 + 当前共有 8 个剩余槽位,因此还需要增加一个副本。专家 1 的当前单副本 + 负载 12 最高,所以将其拆成两个负载为 6 的副本。最终专家组为: + + ``[(expert=1, copies=2, load=6),`` + `` (expert=2..7, copies=1, load=9, 8, 7, 5, 4, 3)]`` + + 3. 第一阶段平铺多副本专家 + + 循环游标从 rank 0 开始,将专家 1 的两个副本依次放到 rank 0、1: + + ``placement = [[0, 1], [0, 1], [0], [0]]`` + ``rank_load = [6, 6, 0, 0]`` + ``remaining_slots = [1, 1, 2, 2]`` + + 4. 第二阶段分配单副本专家 + + 专家 2 到 7 已按负载从高到低排列。每次从最小堆中取当前负载最低的 + 未满 rank;负载相同时使用 rank ID 打破平局: + + * 专家 2,负载 9:放到 rank 2,负载变为 ``[6, 6, 9, 0]``; + * 专家 3,负载 8:放到 rank 3,负载变为 ``[6, 6, 9, 8]``; + * 专家 4,负载 7:放到 rank 0,负载变为 ``[13, 6, 9, 8]``,rank 0 填满; + * 专家 5,负载 5:放到 rank 1,负载变为 ``[13, 11, 9, 8]``,rank 1 填满; + * 专家 6,负载 4:放到 rank 3,负载变为 ``[13, 11, 9, 12]``,rank 3 填满; + * 专家 7,负载 3:放到 rank 2,负载变为 ``[13, 11, 12, 12]``,rank 2 填满。 + + 最终候选布局为: + + ``rank 0: [0, 1, 4]`` + ``rank 1: [0, 1, 5]`` + ``rank 2: [0, 2, 7]`` + ``rank 3: [0, 3, 6]`` + + 将专家 0 的公共负载 10 加回来后,完整 rank 负载为 + ``[23, 21, 22, 22]``,平均负载为 22,最大负载为 23。 + + 5. 复用当前布局 + + 上述 rank 编号只是负载分组结果。规划器会依次处理当前 rank 0 到 3, + 每次从尚未匹配的候选行中选择共同专家最多的一行;匹配完成后,再让 + 共同专家尽量保留原物理槽位。两次排序都不改变负载,完成后直接返回 + 新布局。 + + 合法性和确定性 + -------------- + * 专家覆盖:全卡冗余专家和非冗余专家组来自互斥集合,并且两者合起来 + 包含所有逻辑专家,所以不会遗漏任何专家。 + * 本地去重:一个多副本专家最多有 ``world_size`` 个副本,循环游标在一组 + 副本分配完成前不会第二次经过同一 rank;单副本专家只放置一次。因此 + 同一逻辑专家不会在一个 rank 上出现两次。 + * 槽位守恒:全卡冗余专家放置完成后,剩余副本总数严格等于剩余槽位总数。 + 第一阶段连续平铺使各 rank 的已用槽位数最多相差一个;第二阶段只从尚有 + 空位的 rank 中选择,并在 rank 填满后将其移出最小堆,最终所有槽位恰好 + 填满。 + * 结果确定:专家组排序和最小堆比较最终都使用专家 ID 或 rank ID 打破 + 平局,因此相同输入始终得到相同布局。 """ def __init__( @@ -60,203 +166,202 @@ def plan( logical_expert_load: LogicalExpertLoad, current_placement: ExpertPlacement, ) -> ExpertPlacement: - """根据全局逻辑专家负载生成新的完整布局。 - - ``logical_expert_load`` 的形状为 ``[layer][logical_expert]``,保存所有 - rank 汇总后的负载。``current_placement`` 的形状为 - ``[layer][rank][local_physical_expert]``;每个 rank 的固定主专家在前, - 可迁移冗余专家在后。返回布局与当前布局形状相同。 - """ + """逐层规划专家布局,再组合成完整的多层布局。""" + # 先把外部输入转换成规划器内部统一使用的 Python 数值类型,并一次性 + # 校验所有层的形状和布局约束。后续每层规划之间没有共享的可变状态。 load = [[float(value) for value in layer] for layer in logical_expert_load] current = [[[int(expert) for expert in rank] for rank in layer] for layer in current_placement] self._validate_inputs(load, current) - before_rank_load = self.estimate_rank_load(load, current) - candidates = [self._plan_layer(layer_load, current_layer) for layer_load, current_layer in zip(load, current)] - candidate_rank_load = self.estimate_rank_load(load, candidates) - placement = [] - for layer, candidate in enumerate(candidates): - # 每层独立决定是否采用候选布局。只有最大 rank 负载严格下降时 - # 才迁移,避免相同负载下的无收益调整,也使零负载层保持原布局。 - before = max(before_rank_load[layer]) - after = max(candidate_rank_load[layer]) - improved = candidate != current[layer] and after < before - placement.append(candidate if improved else current[layer]) - return placement - - def estimate_rank_load( - self, - logical_expert_load: LogicalExpertLoad, - placement: ExpertPlacement, - ) -> RankLoad: - """估算每层、每个 rank 上经过对齐后的物理专家工作量。""" - load = [[float(value) for value in layer] for layer in logical_expert_load] - normalized_placement = [[[int(expert) for expert in rank] for rank in layer] for layer in placement] - num_logical_experts = self._validate_inputs(load, normalized_placement) - estimated = [] - for layer_load, layer_placement in zip(load, normalized_placement): - locations = self._expert_locations(layer_placement, num_logical_experts) - rank_load = [0.0] * self.world_size - for expert, expert_locations in enumerate(locations): - for rank, value in enumerate(self._rank_load_for_expert(layer_load[expert], expert_locations)): - rank_load[rank] += value - estimated.append(rank_load) - return estimated + # 每层只依赖自己的逻辑专家负载和当前布局。先完成单层规划,再将结果 + # 按原 layer 顺序组合,避免多层候选和负载数据交叉索引。 + return [self._plan_layer(layer_load, current_layer) for layer_load, current_layer in zip(load, current)] def _plan_layer( self, logical_load: List[float], current_placement: List[List[int]], ) -> List[List[int]]: - num_logical_experts = len(logical_load) - num_primary_experts_per_rank = num_logical_experts // self.world_size - current_redundant_placement = [row[num_primary_experts_per_rank:] for row in current_placement] - - # 从不可变的主专家布局开始。只要目标 rank 尚未持有该专家副本, - # 对应的 (rank, expert) 组合就是合法候选。 - redundant_placement: List[List[int]] = [[] for _ in range(self.world_size)] - locations = [{expert // num_primary_experts_per_rank} for expert in range(num_logical_experts)] - rank_load_by_expert = [ - self._rank_load_for_expert(logical_load[expert], locations[expert]) for expert in range(num_logical_experts) - ] - rank_load = [ - sum(rank_load_by_expert[expert][rank] for expert in range(num_logical_experts)) - for rank in range(self.world_size) + """完成单层副本分配、rank 排布和物理槽位复用。""" + # 阶段 1:选出最热的 R 个专家,并为每个 rank 固定预留它们的副本。 + redundant_experts = self._select_redundant_experts(logical_load) + + # 阶段 2:只在非冗余专家中增加副本,数量恰好填满所有剩余槽位。 + remaining_expert_groups = self._build_remaining_expert_groups(logical_load, redundant_experts) + + # 阶段 3:每个 rank 先放入相同的冗余专家,再通过 rank 优先队列 + # 分配其余专家。同一专家的一组副本会一次性放到不同 rank。 + candidate_placement = self._distribute_remaining_experts(redundant_experts, remaining_expert_groups) + + # 阶段 4:先将候选行贪心匹配到最相似的当前 rank,再复用原物理槽位, + # 依次减少跨 rank 迁移和 rank 内部的槽位搬运。 + return self._reuse_current_slots(candidate_placement, current_placement) + + def _select_redundant_experts(self, logical_load: List[float]) -> List[int]: + """选择需要在所有 rank 上固定放置的最热专家。""" + return sorted(range(len(logical_load)), key=lambda expert: (-logical_load[expert], expert))[ + : self.num_redundant_experts_per_rank ] - total_redundant_experts = self.world_size * self.num_redundant_experts_per_rank - for _ in range(total_redundant_experts): - best_score = None - best_rank = None - best_expert = None - best_rank_load = None - best_expert_rank_load = None - for rank in range(self.world_size): - if len(redundant_placement[rank]) == self.num_redundant_experts_per_rank: - continue + def _build_remaining_expert_groups( + self, + logical_load: List[float], + redundant_experts: List[int], + ) -> List[ExpertReplicaGroup]: + """确定非冗余专家的副本数,并按安全的分配顺序组成专家组。""" + redundant_expert_set = set(redundant_experts) + remaining_experts = [expert for expert in range(len(logical_load)) if expert not in redundant_expert_set] + replica_counts = {expert: 1 for expert in remaining_experts} - for expert in range(num_logical_experts): - if rank in locations[expert]: - continue - - next_locations = locations[expert] | {rank} - next_expert_rank_load = self._rank_load_for_expert(logical_load[expert], next_locations) - next_rank_load = [ - load - rank_load_by_expert[expert][target_rank] + next_expert_rank_load[target_rank] - for target_rank, load in enumerate(rank_load) - ] - - # 首先最小化最大 rank 负载;负载相同时依次选择总对齐工作量 - # 更小、能保留现有本地副本的候选。最后使用 rank/expert ID - # 打破平局,保证相同输入始终得到相同结果。 - score = ( - max(next_rank_load), - sum(next_rank_load), - expert not in current_redundant_placement[rank], - rank, - expert, - ) - if best_score is None or score < best_score: - best_score = score - best_rank = rank - best_expert = expert - best_rank_load = next_rank_load - best_expert_rank_load = next_expert_rank_load - - # 输入校验已保证每个 rank 都有足够多且互不重复的非本地主专家, - # 因此所有冗余槽位一定可以填满。 - assert best_rank is not None and best_expert is not None - assert best_rank_load is not None and best_expert_rank_load is not None - redundant_placement[best_rank].append(best_expert) - locations[best_expert].add(best_rank) - rank_load = best_rank_load - rank_load_by_expert[best_expert] = best_expert_rank_load - - # 布局质量只取决于选中了哪些专家,与它们在本地冗余槽中的顺序无关。 - # 已经选中的现有专家继续使用原槽位,只有新增专家才填入剩余槽位; - # 这样无需干扰上面的负载均衡循环,也能尽量减少权重传输。 - for rank, selected_experts in enumerate(redundant_placement): - selected_set = set(selected_experts) - new_experts = iter(expert for expert in selected_experts if expert not in current_redundant_placement[rank]) - redundant_placement[rank] = [ - expert if expert in selected_set else next(new_experts) for expert in current_redundant_placement[rank] - ] - - return [ - list( - range( - rank * num_primary_experts_per_rank, - (rank + 1) * num_primary_experts_per_rank, - ) + # 每个 rank 的 R 个槽位已由全卡冗余专家占据。此时剩余槽位总数为 + # logical_expert_count,而未分配专家只有 logical_expert_count-R 个, + # 所以还需在非冗余专家中增加 R 个副本。 + for _ in range(self.num_redundant_experts_per_rank): + expert = min( + (expert for expert in remaining_experts if replica_counts[expert] < self.world_size), + key=lambda expert: (-logical_load[expert] / replica_counts[expert], expert), ) - + redundant_experts - for rank, redundant_experts in enumerate(redundant_placement) - ] + replica_counts[expert] += 1 + + # 分配器按专家组工作,而不是把同一专家拆成多个独立元素。这样分配 + # 一组副本时可以暂时取出多个不同 rank,从结构上避免本地重复专家。 + expert_groups = [] + for expert in remaining_experts: + replica_count = replica_counts[expert] + load_per_replica = logical_load[expert] / replica_count + aligned_load_per_replica = ceil(load_per_replica / self.expert_alignment) * self.expert_alignment + expert_groups.append((expert, replica_count, aligned_load_per_replica)) - def _rank_load_for_expert( + # 多副本专家需要在第一阶段先完成平铺,因此排在单副本专家之前。 + # 副本数相同时优先处理单副本负载较高的专家,最后用专家 ID 打破平局。 + expert_groups.sort(key=lambda group: (-group[1], -group[2], group[0])) + return expert_groups + + def _distribute_remaining_experts( self, - logical_expert_load: float, - locations: set, - ) -> List[float]: - """返回该专家在各 rank 上经过对齐后的负载贡献。""" - physical_expert_load = logical_expert_load / len(locations) - aligned_load = ceil(physical_expert_load / self.expert_alignment) * self.expert_alignment - result = [0.0] * self.world_size - for rank in locations: - result[rank] = aligned_load - return result - - def _expert_locations( + redundant_experts: List[int], + expert_groups: List[ExpertReplicaGroup], + ) -> List[List[int]]: + """先平铺多副本专家,再按当前 rank 负载分配单副本专家。""" + placement = [list(redundant_experts) for _ in range(self.world_size)] + + total_replica_count = sum(replica_count for _, replica_count, _ in expert_groups) + assert total_replica_count % self.world_size == 0 + remaining_slots_per_rank = total_replica_count // self.world_size + remaining_slots = [remaining_slots_per_rank] * self.world_size + rank_load = [0.0] * self.world_size + + replicated_expert_groups = [group for group in expert_groups if group[1] > 1] + single_expert_groups = [group for group in expert_groups if group[1] == 1] + + # 阶段 1:用同一个循环游标依次平铺所有多副本专家。每个专家最多有 + # world_size 个副本,所以一组副本在游标绕回起点之前已经分配完毕, + # 同一 rank 不会出现该专家的两个副本。连续使用同一个游标还会让 + # 各 rank 在第一阶段获得的槽位数最多相差一个,不会提前填满某个 rank。 + next_rank = 0 + for expert, replica_count, load_per_replica in replicated_expert_groups: + for _ in range(replica_count): + assert remaining_slots[next_rank] > 0 + placement[next_rank].append(expert) + remaining_slots[next_rank] -= 1 + rank_load[next_rank] += load_per_replica + next_rank = (next_rank + 1) % self.world_size + + # 阶段 2:多副本专家的位置固定后,再把尚有空位的 rank 按当前负载 + # 放入最小堆。这里负载是第一优先级,剩余槽位数不再参与排序;每次 + # 都把当前最热的单副本专家交给最轻的未满 rank。 + # 全卡冗余专家对每个 rank 的贡献相同,因此无需计入 rank_load。 + rank_queue = [(rank_load[rank], rank) for rank in range(self.world_size) if remaining_slots[rank] > 0] + heapq.heapify(rank_queue) + + for expert, replica_count, load_per_replica in single_expert_groups: + assert replica_count == 1 + assert rank_queue, "not enough rank slots to place remaining experts" + + current_load, rank = heapq.heappop(rank_queue) + placement[rank].append(expert) + remaining_slots[rank] -= 1 + + # remaining_slots == 0 表示该 rank 已经刚好填满,不再放回队列。 + if remaining_slots[rank] > 0: + heapq.heappush(rank_queue, (current_load + load_per_replica, rank)) + + assert not rank_queue + assert all(slots == 0 for slots in remaining_slots) + + return placement + + def _reuse_current_slots( self, - placement: List[List[int]], - num_logical_experts: int, - ) -> List[set]: - """反转单层布局,并拒绝同一 rank 上的重复专家副本。""" - locations = [set() for _ in range(num_logical_experts)] - for rank, row in enumerate(placement): - for expert in row: - if rank in locations[expert]: - raise ValueError(f"logical expert {expert} appears twice on rank {rank}") - locations[expert].add(rank) - return locations + candidate_placement: List[List[int]], + current_placement: List[List[int]], + ) -> List[List[int]]: + """贪心匹配候选 rank,并让共同专家尽量复用当前物理槽位。""" + # 步骤 1:准备所有尚未匹配的候选行。 + # + # 候选布局中的 rank 编号只是负载规划阶段产生的临时编号。任意交换 + # 两个候选行都不会改变每行的专家组合和整体负载,因此可以重新排列 + # 候选行,使其尽量贴近当前运行布局。这里同时缓存专家集合,后续可以 + # 直接用集合交集计算两个 rank 之间的相似度。 + unmatched_candidates = [ + (candidate_rank, candidate_experts, set(candidate_experts)) + for candidate_rank, candidate_experts in enumerate(candidate_placement) + ] + placement = [] + + # 步骤 2:按当前 rank 0 -> N-1 的顺序贪心匹配候选行。 + # + # 相似度定义为两个 rank 共同持有的专家数量。共同专家越多,需要跨 + # rank 传输的专家权重就越少。一个候选行被选中后立即从待选列表移除, + # 从而建立当前 rank 和候选行之间的一一对应关系。 + for current_experts in current_placement: + current_expert_set = set(current_experts) + + # max() 首先选择共同专家数量最多的候选行。相似度相同时,负的 + # candidate_rank 让原候选 rank ID 更小的行优先,保证结果确定。 + best_candidate_index = max( + range(len(unmatched_candidates)), + key=lambda index: ( + len(current_expert_set & unmatched_candidates[index][2]), + -unmatched_candidates[index][0], + ), + ) + _, selected_experts, selected_expert_set = unmatched_candidates.pop(best_candidate_index) + + # 步骤 3:在已经匹配的 rank 内复用当前物理槽位。 + # + # new_experts 只包含候选行新引入的专家,并保持候选行中的原始顺序。 + # 它的数量必然等于当前行中需要被替换的专家数量。 + new_experts = [expert for expert in selected_experts if expert not in current_expert_set] + new_expert_index = 0 + rank_placement = [] + + # 依次检查当前物理槽位:如果槽位中的专家仍被候选行选中,就原地 + # 保留;否则用下一个新专家填充。这样只有真正变化的槽位需要搬运 + # 权重,共同专家不会因为候选行内部顺序不同而发生无意义移动。 + for current_expert in current_experts: + if current_expert in selected_expert_set: + rank_placement.append(current_expert) + continue + + rank_placement.append(new_experts[new_expert_index]) + new_expert_index += 1 + + # 所有需要替换的槽位都应恰好消费一个新专家。 + assert new_expert_index == len(new_experts) + placement.append(rank_placement) + + # 每个当前 rank 都必须匹配且只匹配一个候选行。 + assert not unmatched_candidates + return placement def _validate_inputs( self, logical_expert_load: LogicalExpertLoad, placement: ExpertPlacement, - ) -> int: - """校验完整布局约束,并返回逻辑专家数量。""" + ) -> None: + """拒绝会导致逐层规划静默截断的输入。""" if not logical_expert_load: raise ValueError("logical_expert_load must contain at least one layer") if len(placement) != len(logical_expert_load): raise ValueError("load and placement must have the same number of layers") - num_logical_experts = len(logical_expert_load[0]) - if num_logical_experts == 0 or num_logical_experts % self.world_size: - raise ValueError("logical expert count must be positive and divisible by world_size") - if self.num_redundant_experts_per_rank > num_logical_experts - num_logical_experts // self.world_size: - raise ValueError("too many redundant slots to avoid local or duplicate replicas") - num_primary_experts_per_rank = num_logical_experts // self.world_size - num_local_experts_per_rank = num_primary_experts_per_rank + self.num_redundant_experts_per_rank - - for layer_load, layer_placement in zip(logical_expert_load, placement): - if len(layer_load) != num_logical_experts: - raise ValueError("each load layer must contain every logical expert") - if any(value < 0 for value in layer_load): - raise ValueError("logical expert load must be non-negative") - if len(layer_placement) != self.world_size or any( - len(rank) != num_local_experts_per_rank for rank in layer_placement - ): - raise ValueError("each placement layer must be [world_size][local_experts]") - if any(expert < 0 or expert >= num_logical_experts for rank in layer_placement for expert in rank): - raise ValueError("placement contains an invalid logical expert") - for rank, local_expert_ids in enumerate(layer_placement): - expected_primary_expert_ids = list( - range( - rank * num_primary_experts_per_rank, - (rank + 1) * num_primary_experts_per_rank, - ) - ) - if local_expert_ids[:num_primary_experts_per_rank] != expected_primary_expert_ids: - raise ValueError("placement primary experts do not match their owning rank") - self._expert_locations(layer_placement, num_logical_experts) - return num_logical_experts 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 1122d06058..714c6d827b 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 @@ -38,9 +38,9 @@ def __init__(self, *args, **kwargs): def _init_eplb_runtime(self): """初始化本地物理槽位以及可更新的 EPLB 路由运行态。 - ``local_logics_expert_ids_list`` 始终描述全部本地物理行:固定主专家行 - 在前,冗余专家行在后。负载均衡只允许替换冗余后缀,并在同一个安全 - 推理边界同时更新专家权重行和 ``logical_to_physical_map``。 + ``local_logics_expert_ids_list`` 始终描述全部本地物理行。初始化时主专家 + 在前、冗余专家在后;负载均衡运行后允许替换任意物理行,并在同一个 + 安全推理边界同时更新专家权重和 ``logical_to_physical_map``。 """ world_size = get_global_world_size() assert self.n_routed_experts % world_size == 0 diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index b90df92e70..2e5a73acb2 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -75,7 +75,6 @@ def __init__(self, model: TpPartBaseModel) -> None: first_impl = self._eplb_impls[0] self.num_logical_experts: int = first_impl.n_routed_experts self.num_redundant_experts_per_rank: int = first_impl.num_redundant_experts_per_rank - self.num_primary_experts_per_rank: int = self.num_logical_experts // self.world_size # 评估调度:steps 只在 COLLECTING 状态递增。route counter 从当前 # 布局生效时开始累计,让低流量服务可以跨多个评估周期收集足够样本。 @@ -86,7 +85,8 @@ def __init__(self, model: TpPartBaseModel) -> None: self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") - # 每层布局都保存完整的本地专家列表:固定主专家在前,冗余专家在后。 + # 每层布局都保存完整的本地专家列表;完成初始化后,所有物理槽位 + # 都可以由 EPLB 重新分配,不再区分固定主专家槽和冗余专家槽。 # 本 rank 的布局索引为 [layer][local_expert]。 local_expert_ids_by_layer = [list(impl.local_logics_expert_ids_list) for impl in self._eplb_impls] @@ -240,12 +240,12 @@ def _step_wait_plan_finish(self) -> None: self.target_placement: ExpertPlacement = [ [list(expert_ids) for expert_ids in layer_placement] for layer_placement in placement ] - self.pending_transfer_infos = [ - transfer_info + self.pending_transfer_batches = [ + transfer_batch for layer_index, (current_layer, target_layer) in enumerate( zip(self.current_placement, self.target_placement) ) - for transfer_info in build_transfer_plan( + for transfer_batch in build_transfer_plan( current_layer, target_layer, layer_index, @@ -253,7 +253,7 @@ def _step_wait_plan_finish(self) -> None: self.world_size, ) ] - if not self.pending_transfer_infos: + if not self.pending_transfer_batches: raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") self.state = EPLBManagerState.TRANSFERRING if self.global_rank == 0: @@ -264,7 +264,7 @@ def _step_wait_plan_finish(self) -> None: "eplb started steps=%s changed_layer_count=%s changed_slot_count=%s", self.steps, changed_layer_count, - len(self.pending_transfer_infos), + sum(len(transfer_batch) for transfer_batch in self.pending_transfer_batches), ) def _step_transferring(self) -> None: @@ -274,14 +274,14 @@ def _step_transferring(self) -> None: # 没有活动批次时,所有 rank 根据相同的 pending 列表构造下一批任务。 if not hasattr(self, "active_transfer_batch"): - transfer_batch = self._pop_next_transfer_batch() + transfer_batch = self.pending_transfer_batches.pop(0) if self.pending_transfer_batches else [] # 空批次表示公共任务列表已经耗尽,所有 rank 可以同时结束重排。 if not transfer_batch: self.current_placement = self.target_placement elapsed = time.time() - self.rebalance_started_at self._clear_route_counters() - del self.pending_transfer_infos + del self.pending_transfer_batches del self.target_placement del self.rebalance_started_at self.state = EPLBManagerState.COLLECTING @@ -291,65 +291,43 @@ def _step_transferring(self) -> None: self.active_transfer_batch = transfer_batch - # 一个批次内每个 rank 至多参与一条任务。参与者构造并启动本地传输; - # 其他 rank 只保存相同的批次信息,后续共同参与状态同步和 commit。 - local_transfer_info = next( - ( - transfer_info - for transfer_info in transfer_batch - if self.global_rank in (transfer_info.source_rank, transfer_info.dest_rank) - ), - None, - ) - if local_transfer_info is not None: - self.active_transfer = PinnedMemoryEPLBTransfer( + # 普通批次只有一个任务;覆盖环批次可能要求同一 rank 同时保存 + # 多个源/目标的 pinned row,必须等整批传输完成后再统一覆盖 live 权重。 + self.active_transfers = [ + PinnedMemoryEPLBTransfer( self._weights, self.transfer_group, self.global_rank, - local_transfer_info, + transfer_info, ) - self.active_transfer.start() + for transfer_info in transfer_batch + if self.global_rank in (transfer_info.source_rank, transfer_info.dest_rank) + ] + for transfer in self.active_transfers: + transfer.start() else: # 已有活动批次时,本 step 只负责轮询;整批完成后才统一提交。 self._poll_transfer_batch() - def _pop_next_transfer_batch(self) -> List[EPLBTransferInfo]: - """从 pending 中取出一组 rank 互不冲突的传输任务。""" - occupied_ranks = set() - transfer_batch: List[EPLBTransferInfo] = [] - remaining_transfer_infos: List[EPLBTransferInfo] = [] - - # 所有 rank 持有相同的计划列表,并执行相同的贪心扫描,因此会得到完全 - # 相同的批次。每个 rank 在本批次中至多参与一条任务,可以独立启动。 - for transfer_info in self.pending_transfer_infos: - participant_ranks = {transfer_info.source_rank, transfer_info.dest_rank} - # 与已选任务没有公共 rank 时,当前任务可以并入本批次。 - if not participant_ranks & occupied_ranks: - transfer_batch.append(transfer_info) - occupied_ranks |= participant_ranks - else: - remaining_transfer_infos.append(transfer_info) - - # 选中的任务由 active_transfer_batch 持有;未选中的冲突任务保留到下一批。 - self.pending_transfer_infos = remaining_transfer_infos - return transfer_batch - def _poll_transfer_batch(self) -> None: """等待当前批次全部完成,随后统一提交并释放本地任务。""" - active_transfer = getattr(self, "active_transfer", None) - local_state: Optional[Tuple[EPLBTransferInfo, bool]] = None - if active_transfer is not None: - local_state = ( - active_transfer.transfer_info, - active_transfer.is_finished(), - ) - transfer_states: List[Optional[Tuple[EPLBTransferInfo, bool]]] = [None] * self.world_size - dist.all_gather_object(transfer_states, local_state, group=self.control_group) + active_transfers = self.active_transfers + local_states = [(transfer.transfer_info, transfer.is_finished()) for transfer in active_transfers] + transfer_states: List[List[Tuple[EPLBTransferInfo, bool]]] = [[] for _ in range(self.world_size)] + dist.all_gather_object(transfer_states, local_states, group=self.control_group) # 批次内任意任务只要缺少参与方状态,或任一参与方尚未完成,整批都不能 - # commit。下一次 step 会继续轮询同一个批次。 - if not all(state is None or state[1] for state in transfer_states): - return + # commit。跨 rank 任务应收到 source/destination 两份状态,本地复制只需一份。 + for transfer_info in self.active_transfer_batch: + participant_states = [ + finished + for rank_states in transfer_states + for reported_info, finished in rank_states + if reported_info == transfer_info + ] + expected_participant_count = 1 if transfer_info.source_rank == transfer_info.dest_rank else 2 + if len(participant_states) != expected_participant_count or not all(participant_states): + return # 所有 rank 使用相同的批次顺序提交,因此全局 placement 和 metadata # 始终一致;只有 destination rank 会额外写入实际专家权重。 @@ -358,19 +336,21 @@ def _poll_transfer_batch(self) -> None: torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) for transfer_info in self.active_transfer_batch: self._commit_transfer(transfer_info) + for layer_index in {transfer_info.layer_index for transfer_info in self.active_transfer_batch}: + self._publish_layer_metadata(layer_index) - if active_transfer is not None: - del self.active_transfer + del self.active_transfers del self.active_transfer_batch def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: - """提交一条传输,并发布更新后的路由 metadata。""" + """把一条已完成传输提交到 live 权重和完整布局。""" is_destination_rank = transfer_info.dest_rank == self.global_rank if is_destination_rank: - active_transfer = getattr(self, "active_transfer", None) - assert ( - active_transfer is not None and active_transfer.transfer_info == transfer_info - ), "EPLB destination rank has no matching completed transfer" + active_transfer = next( + (transfer for transfer in self.active_transfers if transfer.transfer_info == transfer_info), + None, + ) + assert active_transfer is not None, "EPLB destination rank has no matching completed transfer" for tensor_buffer in active_transfer.tensor_buffers: tensor_buffer.live_tensor[transfer_info.dest_local_expert_index].copy_(tensor_buffer.pinned_row) @@ -384,6 +364,9 @@ def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: transfer_info.dest_local_expert_index ] = transfer_info.source_logical_expert_id + def _publish_layer_metadata(self, layer_index: int) -> None: + """在整批槽位更新完成后发布该层路由 metadata。""" + layer_impl = self._eplb_impls[layer_index] logical_to_physical_map = torch.tensor( build_logical_to_physical_map( self.current_placement[layer_index], diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index b7498051a5..de8f4ce04f 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -194,49 +194,199 @@ def build_transfer_plan( layer_index: int, num_logical_experts: int, world_size: int, -) -> List[EPLBTransferInfo]: - """生成一层中所有发生变化的冗余专家传输任务。 +) -> List[List[EPLBTransferInfo]]: + """生成一层中所有发生变化的专家传输批次。 ``current_placement`` 和 ``target_placement`` 的形状均为 - ``[world_size, num_local_experts]``,每行固定主专家在前、冗余专家在后。 - 每个逻辑专家的固定主副本始终作为传输源,不从已有冗余副本中选择数据源。 + ``[world_size, num_local_experts]``。所有物理槽位都允许变化,因此传输源 + 必须从当前实际存在的副本中选择。 + + 规划分为三个阶段: + + 1. 建立当前槽位索引,同时找出布局调整前后专家不变的稳定槽位; + 2. 为每个变化的目标槽位绑定一个确定的源槽位。优先使用稳定副本,因为 + 这种源槽位永远不会被本轮迁移覆盖;没有稳定副本时,循环使用当前已有 + 的各个副本,避免把全部读取集中到同一个 rank; + 3. 根据槽位覆盖依赖生成两类执行批次。目标槽位不再作为任何待处理任务源 + 的任务属于“安全任务”,彼此不要求原子提交,可以继续拆成较小批次并 + 提前 commit;若不存在安全任务,剩余依赖必然由一个或多个环组成,环内 + 任务不可拆分,必须全部传入 pinned memory 后统一 commit。 + + 图中使用 ``[槽位:当前专家] --传输专家--> [目标槽位]`` 表示一条任务。 + + 链式依赖示例 + ------------ + 当前布局和目标布局分别为: + + ``current: [A:e0] [B:e1] [C:e2] [D:e2]`` + ``target: [A:e0] [B:e0] [C:e1] [D:e2]`` + + 需要执行的传输形成一条依赖链: + + ``[A:e0] --e0--> [B:e1] --e1--> [C:e2]`` + + 初始 ``source_slots={A, B}``,因此只有 C 可以覆盖。虽然 C 中的 e2 被 + 覆盖,但 D 中仍有稳定的 e2;第一批执行 ``B --e1--> C`` 后,e1 已经在 + C 中建立新副本,B 才不再作为源。第二批再执行 ``A --e0--> B``,最终 + 得到目标布局。 + + 这个过程同时保护传入和被覆盖的专家:如果 B 保存的是 e1 的最后一个在线 + 副本,而目标布局仍要求保留 e1,那么一定存在一条以 B 为源的待处理任务。 + 此时 ``B in source_slots``,任何以 B 为目标的任务都不会进入安全批次。只有 + e1 已经存在于其他稳定槽位,或者前一批已经为 e1 建立新位置后,B 才允许 + 被覆盖。 + + 环形依赖示例 + ------------ + 当前布局为 ``[A:e0] [B:e1] [C:e2]``,目标布局为 + ``[A:e1] [B:e2] [C:e0]``,依赖关系为: + + ``[A:e0] --e0--> [C:e2] --e2--> [B:e1] --e1--> [A:e0]`` + + A、B、C 都既是源又是目标,不存在安全目标槽位。三条任务必须组成同一个 + 批次:先将 e0、e1、e2 全部传入 pinned memory,等全部传输完成后再统一 + 覆盖 A、B、C,最后发布新 metadata。 + + 返回值按执行顺序保存最终的 commit 批次:同一安全波次会按参与 rank 拆成 + 若干小批次,每个 rank 在一个小批次中最多参与一条任务;不冲突的 rank 仍 + 可并行传输,已完成的小批次也可以立即提交。环批次则始终包含完整环。这样 + 普通批次在每个 rank 上最多缓存一个专家,只有环形依赖才需要同时缓存多个 + 专家,同时仍保证任何源专家都不会在最后一次读取之前被覆盖。 """ assert world_size > 0 assert num_logical_experts % world_size == 0 assert len(current_placement) == len(target_placement) == world_size - num_primary_experts_per_rank = num_logical_experts // world_size num_local_experts_per_rank = len(current_placement[0]) - assert num_local_experts_per_rank >= num_primary_experts_per_rank assert all(len(row) == num_local_experts_per_rank for row in current_placement) assert all(len(row) == num_local_experts_per_rank for row in target_placement) - + assert all(0 <= expert < num_logical_experts for row in current_placement for expert in row) + assert all(0 <= expert < num_logical_experts for row in target_placement for expert in row) + + # 阶段 1:记录每个逻辑专家当前位于哪些物理槽位,并单独记录不会变化的 + # 稳定槽位。槽位统一表示为 (rank, local_expert_index)。 + Slot = tuple[int, int] + current_slots_by_expert: List[List[Slot]] = [[] for _ in range(num_logical_experts)] + stable_slots_by_expert: List[List[Slot]] = [[] for _ in range(num_logical_experts)] for rank, (current_row, target_row) in enumerate(zip(current_placement, target_placement)): - expected_primary_experts = list( - range( - rank * num_primary_experts_per_rank, - (rank + 1) * num_primary_experts_per_rank, - ) - ) - assert list(current_row[:num_primary_experts_per_rank]) == expected_primary_experts - assert list(target_row[:num_primary_experts_per_rank]) == expected_primary_experts - - transfer_infos: List[EPLBTransferInfo] = [] + for local_expert_index, (current_expert, target_expert) in enumerate(zip(current_row, target_row)): + slot = (rank, local_expert_index) + current_slots_by_expert[current_expert].append(slot) + if current_expert == target_expert: + stable_slots_by_expert[current_expert].append(slot) + assert all(current_slots_by_expert), "current placement must contain every logical expert" + assert set(expert for row in target_placement for expert in row) == set(range(num_logical_experts)) + + # 阶段 2:为每个变化的目标槽位绑定一个确定的当前源槽位。 + # + # 稳定副本不会出现在任何任务的目标位置,因此可以反复读取而没有覆盖 + # 风险。只有不存在稳定副本时,才循环使用该专家当前已有的所有副本。 + # pending_transfers 的每一项为 (source_slot, transfer_info)。source_slot + # 只用于规划覆盖依赖,真正执行任务所需的信息保存在 transfer_info 中。 + source_use_count = [0] * num_logical_experts + pending_transfers: List[tuple[Slot, EPLBTransferInfo]] = [] for destination_rank, (current_row, target_row) in enumerate(zip(current_placement, target_placement)): - for destination_local_expert_index in range(num_primary_experts_per_rank, num_local_experts_per_rank): + for destination_local_expert_index in range(num_local_experts_per_rank): current_expert_id = current_row[destination_local_expert_index] target_expert_id = target_row[destination_local_expert_index] if target_expert_id == current_expert_id: continue - assert 0 <= target_expert_id < num_logical_experts - source_rank = target_expert_id // num_primary_experts_per_rank - transfer_infos.append( - EPLBTransferInfo( - source_rank=source_rank, - layer_index=layer_index, - source_logical_expert_id=target_expert_id, - dest_rank=destination_rank, - dest_local_expert_index=destination_local_expert_index, - ) + source_slots = stable_slots_by_expert[target_expert_id] or current_slots_by_expert[target_expert_id] + source_slot = source_slots[source_use_count[target_expert_id] % len(source_slots)] + source_use_count[target_expert_id] += 1 + transfer_info = EPLBTransferInfo( + source_rank=source_slot[0], + layer_index=layer_index, + source_logical_expert_id=target_expert_id, + dest_rank=destination_rank, + dest_local_expert_index=destination_local_expert_index, ) - - return transfer_infos + pending_transfers.append((source_slot, transfer_info)) + + # 阶段 3:按照槽位覆盖依赖,将任务拆成可安全提交的执行批次。 + transfer_batches: List[List[EPLBTransferInfo]] = [] + while pending_transfers: + source_slots = {source_slot for source_slot, _ in pending_transfers} + + # 3.1 收集当前拓扑层次的全部安全任务。source_slots 是当前仍需保护的 + # 槽位集合:只要某个槽位中的专家尚未完成最后一次读取,该槽位就仍在 + # 集合中,任何以它为目标的任务都不能提交。这同时保护了目标槽位里即将 + # 被覆盖的旧专家,避免其最后一个在线副本被提前删除。 + # + # 目标槽位不在 source_slots 的任务可以并行传输,并在整批完成后统一 + # 提交。若旧专家仍需迁往其他位置,该目标槽位必然也是相应任务的源, + # 因而不会在本轮被选中;若它不是源,则旧专家已经有其他可用副本。 + # + # 这里不能在找到第一个任务后立即修改 source_slots。只有整批提交并从 + # pending 中移除后,下一层目标槽位才真正变得安全。 + safe_transfer_batch: List[EPLBTransferInfo] = [] + remaining_transfers: List[tuple[Slot, EPLBTransferInfo]] = [] + for source_slot, transfer_info in pending_transfers: + destination_slot = (transfer_info.dest_rank, transfer_info.dest_local_expert_index) + if destination_slot not in source_slots: + safe_transfer_batch.append(transfer_info) + else: + remaining_transfers.append((source_slot, transfer_info)) + + if safe_transfer_batch: + # 安全任务之间没有原子提交要求,但若同一 rank 在一个批次中参与 + # 多条任务,就会同时创建多份专家 pinned buffer。这里按 rank 冲突 + # 继续拆分:每个 rank 在一个小批次中最多参与一条任务,不冲突的 + # rank 仍可并行传输,从而兼顾吞吐和 pinned memory 峰值。 + unbatched_transfers = safe_transfer_batch + while unbatched_transfers: + current_batch: List[EPLBTransferInfo] = [] + occupied_ranks: set[int] = set() + deferred_transfers: List[EPLBTransferInfo] = [] + + # 顺序扫描尚未分组的任务:rank 不冲突的任务进入当前批次, + # 冲突任务留到下一轮。每轮至少取出一个任务,因此一定结束。 + for transfer_info in unbatched_transfers: + participant_ranks = {transfer_info.source_rank, transfer_info.dest_rank} + if participant_ranks & occupied_ranks: + deferred_transfers.append(transfer_info) + else: + current_batch.append(transfer_info) + occupied_ranks.update(participant_ranks) + + transfer_batches.append(current_batch) + unbatched_transfers = deferred_transfers + + pending_transfers = remaining_transfers + else: + # 3.2 没有叶子时,每个目标槽位也一定是某条任务的源槽位。每个目标 + # 槽位只有一条写入任务,因此此时源槽位也不会重复,剩余依赖图必然 + # 分解为若干互不相交的简单环。任选第一条任务的源槽位,沿着 + # source_slot -> destination_slot 追踪,回到起点便得到一个完整环。 + # + # 这里有 N 个互不重复的目标槽位,并且没有安全任务意味着这 N 个 + # 目标都包含在 source_slots 中。source_slots 最多也只有 N 项,因此 + # 它必然恰好有 N 项,即每个源槽位只对应一个目标;先显式校验这个 + # 条件,再构造字典,不会因重复 key 丢失任务。 + # + # 例如 ``S -> A、S -> B、A -> S`` 中,源集合只有 ``{S, A}``, + # B 不在源集合中,所以 ``S -> B`` 会先作为安全任务移除;剩余的 + # ``S -> A、A -> S`` 才会进入这里,并且每个源都只对应一个目标。 + assert len(source_slots) == len(pending_transfers) + transfer_by_source_slot = dict(pending_transfers) + + cycle_start_slot = pending_transfers[0][0] + source_slot = cycle_start_slot + cycle_batch: List[EPLBTransferInfo] = [] + cycle_source_slots: set[Slot] = set() + + # 从任意源槽位出发,当前任务的目标槽位就是下一条任务的源槽位; + # 目标重新回到起点时,一个完整环便已经收集完成。 + while True: + assert source_slot not in cycle_source_slots + cycle_source_slots.add(source_slot) + transfer_info = transfer_by_source_slot[source_slot] + cycle_batch.append(transfer_info) + destination_slot = (transfer_info.dest_rank, transfer_info.dest_local_expert_index) + if destination_slot == cycle_start_slot: + break + source_slot = destination_slot + + transfer_batches.append(cycle_batch) + pending_transfers = [transfer for transfer in pending_transfers if transfer[0] not in cycle_source_slots] + + return transfer_batches diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 06f7722744..c30111aa5b 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -307,38 +307,41 @@ def test_eplb_planner_builds_legal_concrete_slot_layout(): result = planner.plan(load.sum(dim=1).tolist(), current) placement = result[0] - for rank, row in enumerate(placement): - assert row[:2] == list(range(rank * 2, (rank + 1) * 2)) - row = row[2:] + for row in placement: + assert len(row) == 3 assert len(row) == len(set(row)) - assert all(expert // 2 != rank for expert in row) - assert max(map(max, planner.estimate_rank_load(load.sum(dim=1).tolist(), result))) <= max( - map(max, planner.estimate_rank_load(load.sum(dim=1).tolist(), current)) - ) + assert 0 in row + assert set(expert for row in placement for expert in row) == set(range(8)) + assert any(row[:2] != list(range(rank * 2, (rank + 1) * 2)) for rank, row in enumerate(placement)) assert isinstance(result, list) -def test_eplb_planner_estimator_distributes_global_load_across_copies(): - planner = GreedyEPLBPlanner( - 4, - 1, - expert_alignment=128, - ) - placement = [[[0, 1, 2], [2, 3, 4], [4, 5, 6], [6, 7, 0]]] - load = [[100, 200, 300, 400, 500, 600, 700, 800]] +def test_eplb_planner_returns_deterministic_layout_for_zero_load_experts(): + planner = GreedyEPLBPlanner(2, 1) + current = [[[0, 1, 3], [2, 3, 1]]] - predicted = planner.estimate_rank_load(load, placement) + result = planner.plan([[0, 0, 0, 0]], current) - assert predicted == [[640, 1024, 1280, 1408]] + assert result == [[[0, 1, 3], [2, 0, 1]]] -def test_eplb_planner_does_not_move_zero_load_experts(): +def test_eplb_planner_plans_each_layer_independently_then_combines_results(): planner = GreedyEPLBPlanner(2, 1) - current = [[[0, 1, 3], [2, 3, 1]]] + current_layer = [[0, 1, 3], [2, 3, 1]] + current = [[row[:] for row in current_layer], [row[:] for row in current_layer]] - result = planner.plan([[0, 0, 0, 0]], current) + result = planner.plan( + [ + [1000, 1, 1, 1], + [0, 0, 0, 0], + ], + current, + ) - assert result == current + assert result == [ + [[0, 1, 3], [2, 0, 1]], + [[0, 1, 3], [2, 0, 1]], + ] def test_eplb_planner_iteratively_places_hot_expert_on_idle_rank(): @@ -347,8 +350,138 @@ def test_eplb_planner_iteratively_places_hot_expert_on_idle_rank(): result = planner.plan([[1000, 1, 1, 1]], current) - assert result == [[[0, 1, 3], [2, 3, 0]]] - assert planner.estimate_rank_load([[1000, 1, 1, 1]], result) == [[502.0, 502.0]] + assert result == [[[0, 1, 3], [2, 0, 1]]] + + +def test_eplb_planner_repeatedly_splits_the_hottest_remaining_expert(): + planner = GreedyEPLBPlanner(4, 3) + current = _initial_expert_placement(8, 4, 3).unsqueeze(0).tolist() + + result = planner.plan([[1000, 900, 800, 700, 1, 1, 1, 1]], current) + + replica_counts = [sum(expert in row for row in result[0]) for expert in range(8)] + assert replica_counts == [4, 4, 4, 4, 1, 1, 1, 1] + + +def test_eplb_planner_balances_expert_groups_with_equal_replica_counts(): + planner = GreedyEPLBPlanner(4, 1) + + placement = planner._distribute_remaining_experts( + redundant_experts=[0], + expert_groups=[ + (1, 1, 8.0), + (2, 1, 7.0), + (3, 1, 6.0), + (4, 1, 5.0), + (5, 1, 4.0), + (6, 1, 3.0), + (7, 1, 2.0), + (8, 1, 1.0), + ], + ) + + assert placement == [ + [0, 1, 8], + [0, 2, 7], + [0, 3, 6], + [0, 4, 5], + ] + + +def test_eplb_planner_places_replicas_of_one_expert_on_distinct_ranks(): + planner = GreedyEPLBPlanner(2, 1) + + placement = planner._distribute_remaining_experts( + redundant_experts=[0], + expert_groups=[ + (1, 2, 5.0), + (2, 1, 8.0), + (3, 1, 1.0), + ], + ) + + assert placement == [ + [0, 1, 2], + [0, 1, 3], + ] + assert all(len(row) == len(set(row)) for row in placement) + + +def test_eplb_planner_places_single_replicas_by_rank_load_before_free_slots(): + planner = GreedyEPLBPlanner(4, 1) + + placement = planner._distribute_remaining_experts( + redundant_experts=[0], + expert_groups=[ + (1, 3, 10.0), + (2, 2, 1.0), + (3, 1, 8.0), + (4, 1, 7.0), + (5, 1, 6.0), + (6, 1, 5.0), + (7, 1, 4.0), + (8, 1, 3.0), + (9, 1, 2.0), + ], + ) + + # 多副本专家平铺后,rank 3 的剩余槽位比 rank 1、2 少,但负载最低; + # 因此它仍连续取得最热的两个单副本专家,并率先填满。 + assert placement == [ + [0, 1, 2, 7], + [0, 1, 5, 9], + [0, 1, 6, 8], + [0, 2, 3, 4], + ] + + +def test_eplb_planner_matches_documented_two_stage_distribution_example(): + planner = GreedyEPLBPlanner(4, 1) + expert_groups = [ + (1, 2, 6.0), + (2, 1, 9.0), + (3, 1, 8.0), + (4, 1, 7.0), + (5, 1, 5.0), + (6, 1, 4.0), + (7, 1, 3.0), + ] + + placement = planner._distribute_remaining_experts( + redundant_experts=[0], + expert_groups=expert_groups, + ) + + assert placement == [ + [0, 1, 4], + [0, 1, 5], + [0, 2, 7], + [0, 3, 6], + ] + load_per_replica = {expert: load for expert, _, load in expert_groups} + assert [sum(load_per_replica[expert] for expert in row[1:]) for row in placement] == [13.0, 11.0, 12.0, 12.0] + + +def test_eplb_planner_greedily_matches_candidate_ranks_before_reusing_slots(): + planner = GreedyEPLBPlanner(3, 1) + current = [ + [0, 1, 2], + [3, 4, 5], + [6, 7, 8], + ] + candidate = [ + [3, 4, 9], + [6, 7, 10], + [0, 1, 11], + ] + + placement = planner._reuse_current_slots(candidate, current) + + assert placement == [ + [0, 1, 11], + [3, 4, 9], + [6, 7, 10], + ] def test_eplb_planner_keeps_selected_experts_in_their_current_slots(): @@ -357,10 +490,11 @@ def test_eplb_planner_keeps_selected_experts_in_their_current_slots(): result = planner.plan([[50, 98, 54, 6, 34, 66, 63, 52]], current) - # Rank 0 的专家 3 保留在原来的第二个冗余槽位,仅将第一个槽位 - # 从专家 2 替换为专家 4。 - assert current[0][0][2:] == [2, 3] - assert result[0][0][2:] == [4, 3] + # 只要专家仍分配在同一个 rank,就保留其原物理槽位。 + for current_row, target_row in zip(current[0], result[0]): + for slot, expert in enumerate(current_row): + if expert in target_row: + assert target_row[slot] == expert def test_eplb_planner_fills_every_rank_with_distinct_nonlocal_experts(): @@ -380,9 +514,9 @@ def test_eplb_planner_fills_every_rank_with_distinct_nonlocal_experts(): assert len(result) == len(current) assert all(len(actual) == len(expected) for actual, expected in zip(result[0], current[0])) - for rank, row in enumerate(result[0]): - assert row[:4] == list(range(rank * 4, (rank + 1) * 4)) - assert all(expert // 4 != rank for expert in row[4:]) + for row in result[0]: + assert len(row) == len(set(row)) == 5 + assert set(expert for row in result[0] for expert in row) == set(range(16)) def test_eplb_planner_supports_multiple_redundant_experts_per_rank(): @@ -394,11 +528,11 @@ def test_eplb_planner_supports_multiple_redundant_experts_per_rank(): result = planner.plan(load, current) - for rank, row in enumerate(result[0]): - assert row[:4] == list(range(rank * 4, (rank + 1) * 4)) - row = row[4:] - assert len(row) == len(set(row)) == 3 - assert all(expert // 4 != rank for expert in row) + for row in result[0]: + assert len(row) == len(set(row)) == 7 + replica_counts = [sum(expert in row for row in result[0]) for expert in range(16)] + assert replica_counts[1] == replica_counts[6] == replica_counts[11] == 4 + assert sum(replica_counts) == 28 def test_fused_moe_loads_default_replicas_into_their_physical_rows(): @@ -569,9 +703,10 @@ def test_transfer_plan_respects_explicit_target_slots(): target = [[0, 1, 5, 4], [2, 3, 7, 6], [4, 5, 1, 0], [6, 7, 3, 2]] plan = build_transfer_plan(current, target, 3, num_logical_experts=8, world_size=4) + transfer_infos = [transfer_info for transfer_batch in plan for transfer_info in transfer_batch] - assert all(info.layer_index == 3 for info in plan) - assert {(info.dest_rank, info.source_logical_expert_id) for info in plan} == { + assert all(info.layer_index == 3 for info in transfer_infos) + assert {(info.dest_rank, info.source_logical_expert_id) for info in transfer_infos} == { (rank, target[rank][slot]) for rank in range(4) for slot in range(2, 4) } @@ -1120,25 +1255,56 @@ def fused(**kwargs): assert all(call["w13"] is w13 and call["w2"] is w2 for call in captured) -def test_transfer_plan_always_uses_primary_expert_rank(): +def test_transfer_plan_uses_stable_current_expert_source(): current = [[0, 1, 4, 5], [2, 3, 6, 7], [4, 5, 0, 1], [6, 7, 2, 3]] target = [[0, 1, 6, 5], [2, 3, 6, 7], [4, 5, 0, 4], [6, 7, 2, 3]] plan = build_transfer_plan(current, target, 5, num_logical_experts=8, world_size=4) assert plan == [ - EPLBTransferInfo(3, 5, 6, 0, 2), - EPLBTransferInfo(2, 5, 4, 2, 3), + [ + EPLBTransferInfo(1, 5, 6, 0, 2), + EPLBTransferInfo(2, 5, 4, 2, 3), + ], ] -def test_transfer_plan_uses_same_primary_source_for_repeated_expert(): +def test_transfer_plan_reuses_stable_source_for_repeated_expert(): current = [[0, 1, 0, 1], [2, 3, 2, 3], [4, 5, 4, 5], [6, 7, 4, 7]] target = [[0, 1, 4, 4], [2, 3, 2, 3], [4, 5, 4, 5], [6, 7, 4, 7]] first = build_transfer_plan(current, target, 5, 8, 4) second = build_transfer_plan(current, target, 5, 8, 4) assert first == second assert first == [ - EPLBTransferInfo(2, 5, 4, 0, 2), - EPLBTransferInfo(2, 5, 4, 0, 3), + [EPLBTransferInfo(2, 5, 4, 0, 2)], + [EPLBTransferInfo(2, 5, 4, 0, 3)], + ] + + +def test_transfer_plan_keeps_primary_slot_swap_in_one_atomic_batch(): + current = [[0, 1], [2, 3]] + target = [[2, 1], [0, 3]] + + plan = build_transfer_plan(current, target, 0, num_logical_experts=4, world_size=2) + + assert plan == [ + [ + EPLBTransferInfo(1, 0, 2, 0, 0), + EPLBTransferInfo(0, 0, 0, 1, 0), + ] + ] + + +def test_transfer_plan_keeps_three_way_cycle_in_one_atomic_batch(): + current = [[0], [1], [2]] + target = [[1], [2], [0]] + + plan = build_transfer_plan(current, target, 0, num_logical_experts=3, world_size=3) + + assert plan == [ + [ + EPLBTransferInfo(1, 0, 1, 0, 0), + EPLBTransferInfo(0, 0, 0, 2, 0), + EPLBTransferInfo(2, 0, 2, 1, 0), + ] ] @@ -1191,7 +1357,6 @@ def test_manager_commits_transfer_rows_and_metadata(): manager.global_rank = 0 manager.world_size = 2 manager.num_logical_experts = 6 - manager.num_primary_experts_per_rank = 3 manager.target_placement = target_placement manager.current_placement = [[[0, 1, 2, 3, 2], [3, 4, 5, 0, 2]]] manager._eplb_impls = [ @@ -1210,11 +1375,12 @@ def test_manager_commits_transfer_rows_and_metadata(): tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -5))], ), ] - manager.active_transfer = transfers[0] + manager.active_transfers = [transfers[0]] manager._commit_transfer(transfers[0].transfer_info) assert manager.current_placement[0][0] == [0, 1, 2, 4, 2] - manager.active_transfer = transfers[1] + manager.active_transfers = [transfers[1]] manager._commit_transfer(transfers[1].transfer_info) + manager._publish_layer_metadata(0) assert torch.equal(live[:3], original_primary) assert torch.equal(live[3], torch.full((4,), -4)) @@ -1245,7 +1411,7 @@ def is_finished(self): manager.transfer_group = object() manager._weights = [object(), object()] manager.world_size = 4 - manager.pending_transfer_infos = [remote_info, local_info0, local_info1] + manager.pending_transfer_batches = [[remote_info, local_info0], [local_info1]] manager.target_placement = [ [[0, 1, 2], [2, 3, 3], [4, 5, 2], [6, 7, 0]], [[0, 1, 4], [2, 3, 5], [4, 5, 6], [6, 7, 1]], @@ -1259,6 +1425,7 @@ def is_finished(self): committed = [] cleared_route_counters = [] manager._commit_transfer = committed.append + manager._publish_layer_metadata = lambda _layer_index: None manager._clear_route_counters = lambda: cleared_route_counters.append(True) waits = [] overlap_stream = object() @@ -1278,9 +1445,19 @@ def state(transfer_info, finished): return transfer_info, finished gathered_states = [ - [state(remote_info, True), state(local_info0, False), state(remote_info, True), state(local_info0, False)], - [state(remote_info, True), state(local_info0, True), state(remote_info, True), state(local_info0, True)], - [state(local_info1, True), state(local_info1, True), None, None], + [ + [state(remote_info, True)], + [state(local_info0, False)], + [state(remote_info, True)], + [state(local_info0, False)], + ], + [ + [state(remote_info, True)], + [state(local_info0, True)], + [state(remote_info, True)], + [state(local_info0, True)], + ], + [[state(local_info1, True)], [state(local_info1, True)], [], []], ] local_states = [] @@ -1294,15 +1471,15 @@ def all_gather_object(output, local_state, **_kwargs): assert starts == [local_info0] assert committed == [] assert manager.active_transfer_batch == [remote_info, local_info0] - assert manager.pending_transfer_infos == [local_info1] + assert manager.pending_transfer_batches == [[local_info1]] manager._step_transferring() assert committed == [] - manager.active_transfer.finished = True + manager.active_transfers[0].finished = True manager._step_transferring() assert committed == [remote_info, local_info0] - assert not hasattr(manager, "active_transfer") + assert not hasattr(manager, "active_transfers") manager._step_transferring() assert starts == [local_info0, local_info1] @@ -1316,14 +1493,14 @@ def all_gather_object(output, local_state, **_kwargs): [[0, 1, 2], [2, 3, 3], [4, 5, 2], [6, 7, 0]], [[0, 1, 4], [2, 3, 5], [4, 5, 6], [6, 7, 1]], ] - assert not hasattr(manager, "pending_transfer_infos") + assert not hasattr(manager, "pending_transfer_batches") assert not hasattr(manager, "target_placement") assert not hasattr(manager, "rebalance_started_at") assert cleared_route_counters == [True] assert local_states == [ - state(local_info0, False), - state(local_info0, True), - state(local_info1, True), + [state(local_info0, False)], + [state(local_info0, True)], + [state(local_info1, True)], ] assert waits == [overlap_stream, overlap_stream] @@ -1373,12 +1550,11 @@ def is_finished(self): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) transfer_info = EPLBTransferInfo(0, 0, 0, 0, 0) transfer = Transfer(live, received, transfer_info) - manager.active_transfer = transfer + manager.active_transfers = [transfer] manager.active_transfer_batch = [transfer_info] manager.control_group = object() manager.world_size = 1 - manager.pending_transfer_infos = [transfer_info] - manager.num_primary_experts_per_rank = 1 + manager.pending_transfer_batches = [] manager.num_logical_experts = 1 manager.global_rank = 0 manager.target_placement = [[[0]]] @@ -1468,7 +1644,7 @@ def all_gather_object(output, local_token_count, **_kwargs): def test_manager_step_uses_explicit_state_instead_of_pending_work(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.TRANSFERRING - manager.pending_transfer_infos = [object()] + manager.pending_transfer_batches = [[object()]] manager._plan_task = object() calls = [] manager._step_transferring = lambda: calls.append("transfer") @@ -1509,7 +1685,7 @@ def test_manager_enters_transferring_state_with_planned_work(monkeypatch): def build_plan(*args): build_calls.append(args) - return [transfer_infos[args[2]]] + return [[transfer_infos[args[2]]]] monkeypatch.setattr(manager_module, "build_transfer_plan", build_plan) monkeypatch.setattr( @@ -1521,7 +1697,7 @@ def build_plan(*args): manager._step_wait_plan_finish() assert manager.state is manager_module.EPLBManagerState.TRANSFERRING - assert manager.pending_transfer_infos == transfer_infos + assert manager.pending_transfer_batches == [[transfer_infos[0]], [transfer_infos[1]]] assert build_calls == [ (manager.current_placement[0], placement[0], 0, manager.num_logical_experts, manager.world_size), (manager.current_placement[1], placement[1], 1, manager.num_logical_experts, manager.world_size), @@ -1899,7 +2075,7 @@ def all_gather_object(output, local_expert_ids_by_layer, group): monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) manager = manager_module.EPLBManager(type("Model", (), {})()) assert not hasattr(manager, "_plan_task") - assert not hasattr(manager, "pending_transfer_infos") + assert not hasattr(manager, "pending_transfer_batches") assert manager.state is manager_module.EPLBManagerState.COLLECTING assert (manager.control_group, manager.transfer_group) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 2 diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 14c352a1b2..7182b01e97 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -78,33 +78,53 @@ def _worker(rank, port): current = [[0, 1, 2], [2, 3, 0]] target = [[0, 1, 3], [2, 3, 1]] for expected_layer in range(2): - transfer_infos = build_transfer_plan( + transfer_batches = build_transfer_plan( current, target, expected_layer, num_logical_experts=4, world_size=2, ) - for transfer_info in transfer_infos: - transfer = PinnedMemoryEPLBTransfer(weights, transfer_group, rank, transfer_info) - assert all(buffer.pinned_row.is_pinned() for buffer in transfer.tensor_buffers) - transfer.start() - _wait_for_transfer(transfer, control_group) - if transfer_info.dest_rank == rank: - expected_expert = 3 if rank == 0 else 1 - expected_w13 = expected_layer * 100 + expected_expert - expected_w2 = expected_w13 + 10 - assert [buffer.name for buffer in transfer.tensor_buffers] == [ - "w13.weight", - "w13.weight_scale", - "w2.weight", - "w2.weight_scale", - ] - pinned_rows = [buffer.pinned_row for buffer in transfer.tensor_buffers] - assert torch.all(pinned_rows[0][0] == expected_w13) - assert torch.all(pinned_rows[1][0] == expected_w13 + 0.5) - assert torch.all(pinned_rows[2][0] == expected_w2) - assert torch.all(pinned_rows[3][0] == expected_w2 + 0.5) + for transfer_batch in transfer_batches: + transfers = [PinnedMemoryEPLBTransfer(weights, transfer_group, rank, info) for info in transfer_batch] + for transfer in transfers: + assert all(buffer.pinned_row.is_pinned() for buffer in transfer.tensor_buffers) + transfer.start() + for transfer in transfers: + _wait_for_transfer(transfer, control_group) + if transfer.transfer_info.dest_rank == rank: + expected_expert = 3 if rank == 0 else 1 + expected_w13 = expected_layer * 100 + expected_expert + expected_w2 = expected_w13 + 10 + assert [buffer.name for buffer in transfer.tensor_buffers] == [ + "w13.weight", + "w13.weight_scale", + "w2.weight", + "w2.weight_scale", + ] + pinned_rows = [buffer.pinned_row for buffer in transfer.tensor_buffers] + assert torch.all(pinned_rows[0][0] == expected_w13) + assert torch.all(pinned_rows[1][0] == expected_w13 + 0.5) + assert torch.all(pinned_rows[2][0] == expected_w2) + assert torch.all(pinned_rows[3][0] == expected_w2 + 0.5) + + # 主槽位互换会形成覆盖环。两个方向必须同时完成 GPU -> pinned memory + # 传输后才能 commit,验证同一 rank 上并发的 send/recv 任务可以正常结束。 + swap_target = [[0, 3, 2], [2, 1, 0]] + swap_plan = build_transfer_plan(current, swap_target, 0, num_logical_experts=4, world_size=2) + assert len(swap_plan) == 1 + swap_infos = swap_plan[0] + assert len(swap_infos) == 2 + swap_transfers = [PinnedMemoryEPLBTransfer(weights, transfer_group, rank, info) for info in swap_infos] + for transfer in swap_transfers: + transfer.start() + for transfer in swap_transfers: + _wait_for_transfer(transfer, control_group) + + for transfer in swap_transfers: + if transfer.transfer_info.dest_rank == rank: + expected_expert = transfer.transfer_info.source_logical_expert_id + assert torch.all(transfer.tensor_buffers[0].pinned_row[0] == expected_expert) dist.destroy_process_group() From 1f6f475b588b65cae268c7337cbbe3b5f8ec02f8 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 20 Sep 2026 08:39:34 +0000 Subject: [PATCH 43/72] simplify EPLB runtime and transfer polling --- .../meta_weights/fused_moe/eplb_placement.py | 29 +++++-- .../meta_weights/fused_moe/impl/__init__.py | 4 +- .../fused_moe/impl/deepgemm_impl.py | 3 +- .../triton_kernel/fused_moe/eplb_topk_ids.py | 2 +- .../model_infer/mode_backend/eplb_manager.py | 53 +++++------- unit_tests/common/fused_moe/test_eplb.py | 84 ++++--------------- .../fused_moe/test_eplb_transfer_gpu.py | 1 - 7 files changed, 67 insertions(+), 109 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py index dbe098d5df..f41b757f29 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -184,16 +184,35 @@ def build_logical_to_physical_maps_for_layers( num_logical_experts: int, current_rank: int, ) -> list[list[list[int]]]: - """逐层构建 CPU list 路由表;设备 Tensor 由调用方在边界处创建。 + """逐层构建完整布局对应的 logical-to-physical 路由表。 - 输入 shape 为 ``[num_layers, num_ranks, num_physical_experts_per_rank]``, - 输出 shape 为 ``[num_layers, num_logical_experts, 2 + routing_slots]``。 + ``rank_to_logic_expert_ids_by_layer`` 使用 EPLB 的完整布局格式: + ``[layer][rank][local physical expert] -> logical expert``。其中每个 layer + 必须包含相同数量的 rank,同一层内每个 rank 必须具有相同数量的物理 + 专家槽位;布局同时包含初始主专家和全部冗余专家。 + + 返回值的 shape 为 ``[layer][logical expert][metadata]``。每个 logical + expert 的 metadata 行格式与 :func:`build_logical_to_physical_map` 一致: + + * 第 0 项是该逻辑专家当前有效的物理副本数; + * 第 1 项标记 ``current_rank`` 是否持有本地副本; + * 第 2 项起保存可参与路由的全局 physical expert ID; + * 如果存在本地副本,本地 physical ID 固定排在第一个路由槽; + * 有效副本之后的固定宽度空槽使用 ``-1`` 填充。 + + 各层相互独立,并按输入中的 layer 顺序逐层调用单层构建函数。这样初始化、 + 在线重排和批量 metadata 发布共享完全相同的副本排序及打包规则,不会出现 + 单层接口和多层接口语义偏差。 + + 本函数只构建普通 CPU list,不创建或搬运 Tensor。调用方需要更新 GPU + 路由 metadata 时,应在通信或提交边界统一转换为对应 dtype/device 的 + ``torch.Tensor``。 """ return [ build_logical_to_physical_map( - rank_to_logic_expert_ids, + layer_placement, num_logical_experts, current_rank=current_rank, ) - for rank_to_logic_expert_ids in rank_to_logic_expert_ids_by_layer + for layer_placement in rank_to_logic_expert_ids_by_layer ] diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 44c3091903..16c6f17576 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -23,10 +23,10 @@ def create_fuse_moe_impl( impl_cls = FuseMoeMarlin else: impl_cls = FuseMoeTriton - kwargs = dict( + + return impl_cls( n_routed_experts=n_routed_experts, num_fused_shared_experts=num_fused_shared_experts, routed_scaling_factor=routed_scaling_factor, quant_method=quant_method, ) - return impl_cls(**kwargs) 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 714c6d827b..92ba5c4677 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 @@ -48,7 +48,6 @@ def _init_eplb_runtime(self): self.num_redundant_experts_per_rank = get_env_start_args().eplb_num_redundant_experts_per_rank if self.num_redundant_experts_per_rank > 0: - self.num_primary_experts_per_rank = self.n_routed_experts // world_size self.num_total_physical_experts = self.n_routed_experts + world_size * self.num_redundant_experts_per_rank initial_local_expert_ids_by_rank = build_initial_local_expert_ids( self.n_routed_experts, @@ -66,6 +65,8 @@ def _init_eplb_runtime(self): ).cuda() # 始终按逻辑专家统计负载,冗余副本不会拆散规划器观察到的负载信号。 self.route_counter = torch.zeros(self.n_routed_experts, dtype=torch.int64, device="cuda") + # 动态 EPLB 默认采集路由负载;以后使用配置文件固定专家布局时, + # 可以关闭该开关,避免执行不再需要的 atomic counter 更新。 self.recording = True else: self.num_total_physical_experts = self.n_routed_experts diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py index 3bde6cfbc3..a6fcb7c63e 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py @@ -81,7 +81,7 @@ def eplb_repair_topk_ids( logical_expert_counter: 每个 logical expert 的累计路由次数,shape 为 ``[num_logical_experts]``。 update_logical_expert_counter: 是否将本次 logical 路由结果累计到 - ``logical_expert_counter``。 + ``logical_expert_counter``。固定布局不需要动态重排时可以关闭。 返回: physical expert ID,shape 为 ``[num_tokens, top_k]``。 diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 2e5a73acb2..5c6f52d0ac 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -1,6 +1,6 @@ from enum import Enum import time -from typing import Dict, List, Optional, Tuple +from typing import Dict, List, Optional import torch import torch.distributed as dist @@ -237,22 +237,23 @@ def _step_wait_plan_finish(self) -> None: self.state = EPLBManagerState.COLLECTING return - self.target_placement: ExpertPlacement = [ - [list(expert_ids) for expert_ids in layer_placement] for layer_placement in placement - ] - self.pending_transfer_batches = [ - transfer_batch - for layer_index, (current_layer, target_layer) in enumerate( - zip(self.current_placement, self.target_placement) - ) - for transfer_batch in build_transfer_plan( + # 广播得到的 placement 已经是规划器新建的完整布局,没有外部持有者会 + # 再修改它,因此可直接保存,不需要逐层深拷贝。 + self.target_placement = placement + + # 每层独立构建有序传输批次,再按 layer 顺序拼接。这样一个批次内只 + # 包含同层任务,提交完成后也只需发布该层的路由 metadata。 + self.pending_transfer_batches: List[List[EPLBTransferInfo]] = [] + layer_placements = zip(self.current_placement, self.target_placement) + for layer_index, (current_layer, target_layer) in enumerate(layer_placements): + layer_transfer_batches = build_transfer_plan( current_layer, target_layer, layer_index, self.num_logical_experts, self.world_size, ) - ] + self.pending_transfer_batches.extend(layer_transfer_batches) if not self.pending_transfer_batches: raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") self.state = EPLBManagerState.TRANSFERRING @@ -291,8 +292,9 @@ def _step_transferring(self) -> None: self.active_transfer_batch = transfer_batch - # 普通批次只有一个任务;覆盖环批次可能要求同一 rank 同时保存 - # 多个源/目标的 pinned row,必须等整批传输完成后再统一覆盖 live 权重。 + # 普通批次允许多个 rank 不冲突的任务并行,但每个 rank 最多参与 + # 一条;覆盖环批次可能要求同一 rank 同时保存多个源/目标的 pinned + # row,必须等整批传输完成后再统一覆盖 live 权重。 self.active_transfers = [ PinnedMemoryEPLBTransfer( self._weights, @@ -311,23 +313,14 @@ def _step_transferring(self) -> None: def _poll_transfer_batch(self) -> None: """等待当前批次全部完成,随后统一提交并释放本地任务。""" - active_transfers = self.active_transfers - local_states = [(transfer.transfer_info, transfer.is_finished()) for transfer in active_transfers] - transfer_states: List[List[Tuple[EPLBTransferInfo, bool]]] = [[] for _ in range(self.world_size)] - dist.all_gather_object(transfer_states, local_states, group=self.control_group) - - # 批次内任意任务只要缺少参与方状态,或任一参与方尚未完成,整批都不能 - # commit。跨 rank 任务应收到 source/destination 两份状态,本地复制只需一份。 - for transfer_info in self.active_transfer_batch: - participant_states = [ - finished - for rank_states in transfer_states - for reported_info, finished in rank_states - if reported_info == transfer_info - ] - expected_participant_count = 1 if transfer_info.source_rank == transfer_info.dest_rank else 2 - if len(participant_states) != expected_participant_count or not all(participant_states): - return + # 每个 rank 只负责自己参与的任务;不参与当前批次的 rank,其本地任务 + # 列表为空,all([]) 自然为 True。所有 rank 汇总一个布尔值即可判断整批 + # 是否完成,无需重复传输并逐条匹配 EPLBTransferInfo。 + local_finished = all(transfer.is_finished() for transfer in self.active_transfers) + finished_by_rank = [False] * self.world_size + dist.all_gather_object(finished_by_rank, local_finished, group=self.control_group) + if not all(finished_by_rank): + return # 所有 rank 使用相同的批次顺序提交,因此全局 placement 和 metadata # 始终一致;只有 destination rank 会额外写入实际专家权重。 diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index c30111aa5b..49216d5f4f 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -70,7 +70,6 @@ def _test_moe_impl( num_redundant_experts_per_rank = 0 return SimpleNamespace( n_routed_experts=num_logical_experts, - num_primary_experts_per_rank=num_logical_experts // world_size, num_total_physical_experts=(num_logical_experts + world_size * num_redundant_experts_per_rank), num_redundant_experts_per_rank=num_redundant_experts_per_rank, local_logics_expert_ids_list=list(range(num_logical_experts // world_size + num_redundant_experts_per_rank)), @@ -82,7 +81,6 @@ def _test_moe_impl( def _set_deepgemm_runtime(impl, runtime): for name in ( - "num_primary_experts_per_rank", "num_total_physical_experts", "num_redundant_experts_per_rank", "logical_to_physical_map", @@ -174,17 +172,6 @@ def _fused_experts( assert seen["fused"]["topk_ids"] == "physical_ids" -def test_deepgemm_runtime_derives_expert_layout(): - runtime = _test_moe_impl( - eplb=True, - num_logical_experts=4, - world_size=2, - num_redundant_experts_per_rank=1, - ) - assert runtime.num_primary_experts_per_rank == 2 - assert runtime.num_total_physical_experts == 6 - - def test_factory_selects_all_paths_without_ep_constructor_state(monkeypatch): plain_quant = SimpleNamespace(method_name="none") marlin_quant = SimpleNamespace(method_name="awq_marlin") @@ -652,50 +639,29 @@ def test_current_rank_stably_moves_all_local_physical_ids_to_front(): assert rank1_map[0] == [4, 1, 3, 0, 1, 5] -@pytest.mark.parametrize( - "current_rank", - [0, 1, 2, 3], -) +@pytest.mark.parametrize("current_rank", [0, 1, 2, 3]) def test_logical_to_physical_maps_for_layers_match_single_layer_api(current_rank): - placements_by_layer = torch.tensor( - [ - [[4, 5], [0, 1], [0, 1], [2, 3]], - [[6, 7], [0, 1], [0, 1], [2, 3]], - [[4, 5], [0, 1], [0, 1], [2, 3]], - ], - dtype=torch.int64, - ) - - rank_to_logic_expert_ids_by_layer = [ - _rank_to_logic_expert_ids(placement.tolist(), 8) for placement in placements_by_layer + placements_by_layer = [ + _rank_to_logic_expert_ids([[4, 5], [0, 1], [0, 1], [2, 3]], 8), + _rank_to_logic_expert_ids([[6, 7], [0, 1], [0, 1], [2, 3]], 8), + _rank_to_logic_expert_ids([[4, 5], [0, 1], [0, 1], [2, 3]], 8), ] + maps_by_layer = build_logical_to_physical_maps_for_layers( - rank_to_logic_expert_ids_by_layer, + placements_by_layer, num_logical_experts=8, current_rank=current_rank, ) - expected_by_layer = [ + expected_maps = [ build_logical_to_physical_map( - rank_to_logic_expert_ids, + layer_placement, num_logical_experts=8, current_rank=current_rank, ) - for rank_to_logic_expert_ids in rank_to_logic_expert_ids_by_layer + for layer_placement in placements_by_layer ] - assert maps_by_layer == expected_by_layer - assert all( - physical_expert_id >= 0 - for logical_map in maps_by_layer - for row in logical_map - for physical_expert_id in row[2 : 2 + row[0]] - ) - assert all( - physical_expert_id == -1 - for logical_map in maps_by_layer - for row in logical_map - for physical_expert_id in row[2 + row[0] :] - ) + assert maps_by_layer == expected_maps def test_transfer_plan_respects_explicit_target_slots(): @@ -722,7 +688,6 @@ def test_manager_evaluating_copies_route_counters_to_cpu_without_modifying_them( _test_moe_impl( eplb=True, route_counter=counter, - recording=False, num_logical_experts=2, world_size=1, ) @@ -1148,7 +1113,6 @@ def cpu_zeros(*shape, **kwargs): monkeypatch.setattr(deepgemm_module.torch, "zeros", cpu_zeros) impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace()) - assert impl.num_primary_experts_per_rank == 2 assert impl.num_redundant_experts_per_rank == 1 assert impl.num_total_physical_experts == 6 assert impl.route_counter.shape == (4,) @@ -1441,23 +1405,10 @@ def is_finished(self): lambda _weights, _group, _rank, transfer_info: Transfer(transfer_info), ) - def state(transfer_info, finished): - return transfer_info, finished - gathered_states = [ - [ - [state(remote_info, True)], - [state(local_info0, False)], - [state(remote_info, True)], - [state(local_info0, False)], - ], - [ - [state(remote_info, True)], - [state(local_info0, True)], - [state(remote_info, True)], - [state(local_info0, True)], - ], - [[state(local_info1, True)], [state(local_info1, True)], [], []], + [True, False, True, False], + [True, True, True, True], + [True, True, True, True], ] local_states = [] @@ -1497,11 +1448,7 @@ def all_gather_object(output, local_state, **_kwargs): assert not hasattr(manager, "target_placement") assert not hasattr(manager, "rebalance_started_at") assert cleared_route_counters == [True] - assert local_states == [ - [state(local_info0, False)], - [state(local_info0, True)], - [state(local_info1, True)], - ] + assert local_states == [False, True, True] assert waits == [overlap_stream, overlap_stream] @@ -2088,7 +2035,6 @@ def all_gather_object(output, local_expert_ids_by_layer, group): assert "planner=GreedyEPLBPlanner" in logs[0] assert weight.fuse_moe_impl.recording assert manager._eplb_impls[0] is weight.fuse_moe_impl - assert not hasattr(weight.fuse_moe_impl, "update_logical_expert_counter") @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 7182b01e97..e52cdf853c 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -29,7 +29,6 @@ def __init__(self, rank, layer_index): self.layer_num_ = layer_index logical_ids = ([0, 1, 2], [2, 3, 0])[rank] self.fuse_moe_impl = SimpleNamespace( - num_primary_experts_per_rank=2, num_redundant_experts_per_rank=1, local_logics_expert_ids_list=list(logical_ids), ) From f958adecaa7d422c0c7e77541fa5dc53c540e058 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 02:03:35 +0000 Subject: [PATCH 44/72] add EPLB rebalance count control --- docs/CN/source/tutorial/api_server_args.rst | 12 +++- docs/EN/source/tutorial/api_server_args.rst | 13 +++- lightllm/server/api_cli.py | 7 +++ lightllm/server/api_start.py | 1 + lightllm/server/core/objs/start_args_type.py | 1 + .../model_infer/mode_backend/base_backend.py | 4 +- .../model_infer/mode_backend/eplb_manager.py | 33 +++++++--- unit_tests/common/fused_moe/test_eplb.py | 63 +++++++++++++++++++ .../fused_moe/test_eplb_transfer_gpu.py | 44 +++++++++++++ unit_tests/server/test_api_start_eplb.py | 20 ++++++ 10 files changed, 187 insertions(+), 11 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index c45d8b9bac..6410394459 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -724,12 +724,22 @@ PD 分离模式参数 EPLB 当前不能与 ``--enable_prefill_cudagraph`` 同时使用,也不支持 SM100 GPU。 同一部署中的所有 rank 和节点必须使用相同的配置值。 +.. option:: --eplb_rebalance_count + + 动态 EPLB 最多成功执行的重排次数,默认值为 ``1``。只有新布局实际发生 + 专家权重迁移并完成提交后才计数;样本不足或规划布局不变不会消耗次数。 + + * ``-1``:不限制次数,持续进行动态重排; + * ``0``:不进行动态重排,仅使用初始化时的冗余布局; + * 正整数:完成指定次数的重排后停止规划。 + 以下示例为每个 EP rank 配置两个冗余专家:: python -m lightllm.server.api_server \ --model_dir /path/to/model \ --enable_ep_moe \ - --eplb_num_redundant_experts_per_rank 2 + --eplb_num_redundant_experts_per_rank 2 \ + --eplb_rebalance_count 1 MTP 多预测参数 -------------- diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index f24d16091f..11c7a6a2ee 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -742,12 +742,23 @@ Expert Parallelism and EPLB Parameters EPLB currently cannot be combined with ``--enable_prefill_cudagraph`` and is not supported on SM100 GPUs. Use the same value on every rank and node in one deployment. +.. option:: --eplb_rebalance_count + + Maximum number of successfully completed dynamic EPLB rebalances. The default is ``1``. A count is consumed + only after a new placement has transferred and committed its expert weights; insufficient samples and unchanged + placements do not consume the limit. + + * ``-1`` keeps dynamic rebalancing enabled indefinitely. + * ``0`` disables dynamic rebalancing, leaving only the initial redundant placement active. + * A positive value stops planning after that many completed rebalances. + Example: enable EPLB with two redundant experts per EP rank:: python -m lightllm.server.api_server \ --model_dir /path/to/model \ --enable_ep_moe \ - --eplb_num_redundant_experts_per_rank 2 + --eplb_num_redundant_experts_per_rank 2 \ + --eplb_rebalance_count 1 MTP Multi-Prediction Parameters ------------------------------- diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 23848c20bd..62a60875b4 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -773,6 +773,13 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: help="""Number of redundant physical experts per EP rank for each MoE layer. Set to 0 to disable EPLB.""", ) + parser.add_argument( + "--eplb_rebalance_count", + type=int, + default=1, + help="""Maximum number of completed EPLB rebalances. -1 means unlimited, + 0 disables dynamic rebalancing, and the default is 1.""", + ) parser.add_argument( "--enable_fused_shared_experts", action="store_true", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 771dc9ed0c..94eb52bfa4 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -162,6 +162,7 @@ def _launch_subprocesses(args: StartArgs): assert ( args.eplb_num_redundant_experts_per_rank >= 0 ), "--eplb_num_redundant_experts_per_rank must be greater than or equal to 0" + assert args.eplb_rebalance_count >= -1, "--eplb_rebalance_count must be greater than or equal to -1" if args.eplb_num_redundant_experts_per_rank > 0: assert args.enable_ep_moe, "EPLB requires --enable_ep_moe" assert not args.enable_prefill_cudagraph, "EPLB does not support --enable_prefill_cudagraph" diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 4527dd6387..7cb317c75b 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -187,6 +187,7 @@ class StartArgs: ) enable_ep_moe: bool = field(default=False) eplb_num_redundant_experts_per_rank: int = field(default=0) + eplb_rebalance_count: int = field(default=1) enable_fused_shared_experts: bool = field(default=False) mtp_mode: Optional[str] = field( default=None, diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 4a72ed66ba..5f600e1359 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -257,10 +257,10 @@ def init_model(self, kvargs): prof_name = f"lightllm-model_backend-node{self.node_rank}_dev{get_current_device_id()}" prof_mode = self.args.enable_profiling self.profiler = ProcessProfiler(mode=prof_mode, name=prof_name, use_multi_thread=True) if prof_mode else None - if self.args.eplb_num_redundant_experts_per_rank > 0: + if self.args.eplb_num_redundant_experts_per_rank > 0 and self.args.eplb_rebalance_count != 0: from lightllm.server.router.model_infer.mode_backend.eplb_manager import EPLBManager - self.eplb_manager = EPLBManager(self.model) + self.eplb_manager = EPLBManager(self.model, self.args.eplb_rebalance_count) # 启动infer_loop_thread, 启动两个线程进行推理,对于具备双batch推理折叠得场景 # 可以降低 cpu overhead,大幅提升gpu得使用率。 diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 5c6f52d0ac..08af99d9fd 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -46,6 +46,7 @@ class EPLBManagerState(Enum): PLANNING = "planning" WAIT_PLAN_FINISH = "wait_plan_finish" TRANSFERRING = "transferring" + FINISHED = "finished" class EPLBManager: @@ -53,17 +54,20 @@ class EPLBManager: 状态循环如下: - ``COLLECTING -> EVALUATING -> PLANNING -> WAIT_PLAN_FINISH -> TRANSFERRING -> COLLECTING`` + ``COLLECTING -> EVALUATING -> PLANNING -> WAIT_PLAN_FINISH -> TRANSFERRING`` 当累计的平均专家 token 数不足时,``EVALUATING`` 会回到 ``COLLECTING``;当规划器认为无需调整布局时,``WAIT_PLAN_FINISH`` 会 - 回到 ``COLLECTING``。每次调用 :meth:`step` 最多推进一个状态,布局 - 规划和权重传输在后台执行,主推理线程负责评估、轮询和提交结果。 + 回到 ``COLLECTING``。完成一次重排后,未达到次数上限则重新进入 + ``COLLECTING``,否则进入不再采样和规划的 ``FINISHED``。每次调用 + :meth:`step` 最多推进一个状态,布局规划和权重传输在后台执行,主推理 + 线程负责评估、轮询和提交结果。 """ - def __init__(self, model: TpPartBaseModel) -> None: + def __init__(self, model: TpPartBaseModel, max_rebalance_count: int = 1) -> None: weights: List[FusedMoeWeight] = _find_fused_moe_weights(model) assert weights, "EPLB requires at least one EP MoE layer" + assert max_rebalance_count == -1 or max_rebalance_count > 0 # 模型与专家拓扑:初始化后保持不变。 self._weights: List[FusedMoeWeight] = weights @@ -80,6 +84,8 @@ def __init__(self, model: TpPartBaseModel) -> None: # 布局生效时开始累计,让低流量服务可以跨多个评估周期收集足够样本。 self.step_interval: int = get_eplb_step_interval() self.steps: int = 0 + self.max_rebalance_count: int = max_rebalance_count + self.completed_rebalance_count: int = 0 # 分布式通信:控制面与权重传输使用独立的通信组。 self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") @@ -118,11 +124,15 @@ def __init__(self, model: TpPartBaseModel) -> None: logger.info( f"eplb enabled layers={len(weights)} num_logical_experts={self.num_logical_experts} " f"num_redundant_experts_per_rank={self.num_redundant_experts_per_rank} " - f"step_interval={self.step_interval} planner={type(self.planner).__name__}" + f"step_interval={self.step_interval} max_rebalance_count={self.max_rebalance_count} " + f"planner={type(self.planner).__name__}" ) def step(self) -> None: """在一个安全的推理边界推进一次状态机。""" + if self.state is EPLBManagerState.FINISHED: + return + if self.state is EPLBManagerState.COLLECTING: self._step_collecting() return @@ -282,12 +292,21 @@ def _step_transferring(self) -> None: self.current_placement = self.target_placement elapsed = time.time() - self.rebalance_started_at self._clear_route_counters() + self.completed_rebalance_count += 1 + reached_rebalance_limit = ( + self.max_rebalance_count != -1 and self.completed_rebalance_count >= self.max_rebalance_count + ) del self.pending_transfer_batches del self.target_placement del self.rebalance_started_at - self.state = EPLBManagerState.COLLECTING + self.state = EPLBManagerState.FINISHED if reached_rebalance_limit else EPLBManagerState.COLLECTING if self.global_rank == 0: - logger.info("eplb completed wall_time=%.2fs", elapsed) + logger.info( + "eplb completed wall_time=%.2fs completed_rebalance_count=%s max_rebalance_count=%s", + elapsed, + self.completed_rebalance_count, + self.max_rebalance_count, + ) return self.active_transfer_batch = transfer_batch diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 49216d5f4f..e77066edf4 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -244,6 +244,10 @@ def test_eplb_redundant_experts_default_to_disabled(): assert parser.parse_args([]).eplb_num_redundant_experts_per_rank == 0 assert parser.parse_args(["--eplb_num_redundant_experts_per_rank", "3"]).eplb_num_redundant_experts_per_rank == 3 assert StartArgs().eplb_num_redundant_experts_per_rank == 0 + assert parser.parse_args([]).eplb_rebalance_count == 1 + assert parser.parse_args(["--eplb_rebalance_count", "-1"]).eplb_rebalance_count == -1 + assert parser.parse_args(["--eplb_rebalance_count", "0"]).eplb_rebalance_count == 0 + assert StartArgs().eplb_rebalance_count == 1 @pytest.mark.parametrize( @@ -1122,6 +1126,29 @@ def cpu_zeros(*shape, **kwargs): assert not hasattr(impl, "expert_parallel_state") +def test_deepgemm_keeps_route_recording_when_rebalance_count_is_zero(monkeypatch): + monkeypatch.setattr( + deepgemm_module, + "get_env_start_args", + lambda: SimpleNamespace( + eplb_num_redundant_experts_per_rank=1, + eplb_rebalance_count=0, + ), + ) + monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) + monkeypatch.setattr( + deepgemm_module.torch, + "zeros", + lambda *shape, **kwargs: torch.full(shape, 0, dtype=kwargs.get("dtype")), + ) + + impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace()) + + assert impl.recording + + def test_eplb_prepare_repairs_logical_ids(monkeypatch): impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) runtime = _test_moe_impl(eplb=True, recording=True) @@ -1386,6 +1413,8 @@ def is_finished(self): ] manager.global_rank = 1 manager.state = manager_module.EPLBManagerState.TRANSFERRING + manager.max_rebalance_count = -1 + manager.completed_rebalance_count = 0 committed = [] cleared_route_counters = [] manager._commit_transfer = committed.append @@ -1440,6 +1469,7 @@ def all_gather_object(output, local_state, **_kwargs): manager._step_transferring() assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert manager.completed_rebalance_count == 1 assert manager.current_placement == [ [[0, 1, 2], [2, 3, 3], [4, 5, 2], [6, 7, 0]], [[0, 1, 4], [2, 3, 5], [4, 5, 6], [6, 7, 1]], @@ -1452,6 +1482,28 @@ def all_gather_object(output, local_state, **_kwargs): assert waits == [overlap_stream, overlap_stream] +def test_manager_finishes_after_reaching_rebalance_limit(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + impls = [SimpleNamespace(recording=True), SimpleNamespace(recording=True)] + target_placement = [[[0, 1], [1, 0]]] + manager.state = manager_module.EPLBManagerState.TRANSFERRING + manager.global_rank = 1 + manager._eplb_impls = impls + manager.current_placement = [[[0, 1], [0, 1]]] + manager.target_placement = target_placement + manager.pending_transfer_batches = [] + manager.max_rebalance_count = 1 + manager.completed_rebalance_count = 0 + manager._clear_route_counters = lambda: None + + manager._step_transferring() + + assert manager.current_placement is target_placement + assert manager.completed_rebalance_count == 1 + assert manager.state is manager_module.EPLBManagerState.FINISHED + assert all(impl.recording for impl in impls) + + def test_wait_plan_finish_broadcasts_pending_status(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.WAIT_PLAN_FINISH @@ -1551,6 +1603,15 @@ def test_manager_step_advances_inflight_transfer(): assert calls == ["transfer"] +def test_finished_manager_step_is_a_noop(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.state = manager_module.EPLBManagerState.FINISHED + + manager.step() + + assert manager.state is manager_module.EPLBManagerState.FINISHED + + def test_manager_evaluates_only_after_entering_evaluating_state(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) route_counter = torch.tensor([1, 2], dtype=torch.int64) @@ -2031,6 +2092,8 @@ def all_gather_object(output, local_expert_ids_by_layer, group): assert manager.metric_client is metric_client assert metric_client_ports == [1234] assert manager.next_evaluation_step == manager.step_interval + assert manager.max_rebalance_count == 1 + assert manager.completed_rebalance_count == 0 assert clear_calls == [manager] assert "planner=GreedyEPLBPlanner" in logs[0] assert weight.fuse_moe_impl.recording diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index e52cdf853c..c0ae0b856f 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -1,5 +1,6 @@ """Multi-GPU correctness test for pinned-memory EPLB transfers.""" +import gc import os import socket import time @@ -65,6 +66,49 @@ def _wait_for_transfer(transfer, control_group): raise TimeoutError("EPLB transfer worker did not finish globally") +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_pytorch_keeps_unreferenced_pinned_source_alive_until_async_copy_finishes(): + """验证 pinned allocator 不会提前复用仍被异步 H2D 读取的内存。 + + copy stream 中先排入一个长任务,确保 H2D copy 在 Python 引用释放时仍未 + 完成。随后删除 pinned tensor 的唯一引用,并立刻申请同尺寸 pinned 内存: + + * 如果旧地址尚未复用,新 buffer 可以立即覆盖而不影响 H2D; + * 如果 allocator 返回了旧地址,对应 copy event 必须已经完成; + * 最终 GPU 数据必须保持为源 buffer 的原始内容。 + + 该测试验证当前 PyTorch/CUDA 组合的运行时行为;allocator 的正确性契约 + 仍由 PyTorch ``copy_`` 中的 host ``record_event`` 实现提供。 + """ + torch.cuda.synchronize() + num_elements = 8 * 1024 * 1024 + source = torch.full((num_elements,), 7, dtype=torch.int32, pin_memory=True) + source_ptr = source.data_ptr() + destination = torch.empty_like(source, device="cuda") + + copy_stream = torch.cuda.Stream() + copy_finished = torch.cuda.Event() + with torch.cuda.stream(copy_stream): + # copy 与 sleep 位于同一 stream,必须等 sleep 完成后才能开始。 + torch.cuda._sleep(1_000_000_000) + destination.copy_(source, non_blocking=True) + copy_finished.record() + + assert not copy_finished.query(), "test setup failed to leave the H2D copy pending" + + del source + gc.collect() + + replacement = torch.empty((num_elements,), dtype=torch.int32, pin_memory=True) + if replacement.data_ptr() == source_ptr: + # 相同地址只有在原 copy 已经结束、allocator 确认可以复用后才合法。 + assert copy_finished.query() + replacement.fill_(-3) + + copy_finished.synchronize() + assert torch.all(destination == 7) + + def _worker(rank, port): os.environ["MASTER_ADDR"] = "127.0.0.1" os.environ["MASTER_PORT"] = str(port) diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py index 36c7fc48f6..a1da20b1ba 100644 --- a/unit_tests/server/test_api_start_eplb.py +++ b/unit_tests/server/test_api_start_eplb.py @@ -24,6 +24,26 @@ def test_eplb_redundant_expert_count_must_not_be_negative(monkeypatch): api_start._launch_subprocesses(args) +def test_eplb_rebalance_count_must_not_be_less_than_negative_one(monkeypatch): + args = StartArgs( + eplb_rebalance_count=-2, + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) + monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) + + with pytest.raises( + AssertionError, + match="--eplb_rebalance_count must be greater than or equal to -1", + ): + api_start._launch_subprocesses(args) + + def test_eplb_redundant_experts_require_ep_moe(monkeypatch): args = StartArgs( enable_ep_moe=False, From a0dd26e8657a490852ce0244b007346a36cb2d21 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 02:20:41 +0000 Subject: [PATCH 45/72] keep reporting EPLB load after rebalancing --- .../model_infer/mode_backend/base_backend.py | 2 +- .../model_infer/mode_backend/eplb_manager.py | 51 +++++++++-------- unit_tests/common/fused_moe/test_eplb.py | 56 +++++++++++++++---- 3 files changed, 72 insertions(+), 37 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 5f600e1359..ec85ef535f 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -257,7 +257,7 @@ def init_model(self, kvargs): prof_name = f"lightllm-model_backend-node{self.node_rank}_dev{get_current_device_id()}" prof_mode = self.args.enable_profiling self.profiler = ProcessProfiler(mode=prof_mode, name=prof_name, use_multi_thread=True) if prof_mode else None - if self.args.eplb_num_redundant_experts_per_rank > 0 and self.args.eplb_rebalance_count != 0: + if self.args.eplb_num_redundant_experts_per_rank > 0: from lightllm.server.router.model_infer.mode_backend.eplb_manager import EPLBManager self.eplb_manager = EPLBManager(self.model, self.args.eplb_rebalance_count) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 08af99d9fd..ca60ec4d80 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -46,7 +46,6 @@ class EPLBManagerState(Enum): PLANNING = "planning" WAIT_PLAN_FINISH = "wait_plan_finish" TRANSFERRING = "transferring" - FINISHED = "finished" class EPLBManager: @@ -58,8 +57,8 @@ class EPLBManager: 当累计的平均专家 token 数不足时,``EVALUATING`` 会回到 ``COLLECTING``;当规划器认为无需调整布局时,``WAIT_PLAN_FINISH`` 会 - 回到 ``COLLECTING``。完成一次重排后,未达到次数上限则重新进入 - ``COLLECTING``,否则进入不再采样和规划的 ``FINISHED``。每次调用 + 回到 ``COLLECTING``。达到重排次数上限后,manager 仍会周期性采集并 + 上报本地负载指标,但不再进入全局负载汇总和布局规划。每次调用 :meth:`step` 最多推进一个状态,布局规划和权重传输在后台执行,主推理 线程负责评估、轮询和提交结果。 """ @@ -67,7 +66,7 @@ class EPLBManager: def __init__(self, model: TpPartBaseModel, max_rebalance_count: int = 1) -> None: weights: List[FusedMoeWeight] = _find_fused_moe_weights(model) assert weights, "EPLB requires at least one EP MoE layer" - assert max_rebalance_count == -1 or max_rebalance_count > 0 + assert max_rebalance_count >= -1 # 模型与专家拓扑:初始化后保持不变。 self._weights: List[FusedMoeWeight] = weights @@ -130,9 +129,6 @@ def __init__(self, model: TpPartBaseModel, max_rebalance_count: int = 1) -> None def step(self) -> None: """在一个安全的推理边界推进一次状态机。""" - if self.state is EPLBManagerState.FINISHED: - return - if self.state is EPLBManagerState.COLLECTING: self._step_collecting() return @@ -167,16 +163,27 @@ def _step_collecting(self) -> None: self.state = EPLBManagerState.EVALUATING def _step_evaluating(self) -> None: - """将负载复制到 CPU,并根据全局样本量进入采样或规划状态。""" + """发布本地负载指标,并在次数允许时根据全局样本量决定是否规划。""" counters = [impl.route_counter for impl in self._eplb_impls] if any(counter.ndim != 1 or counter.shape[0] != self.num_logical_experts for counter in counters): raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") # 将各层累计的路由计数复制到 CPU,后续规划统一使用这份快照。 - # 此处有意不清零 GPU counter:如果样本不足或无需迁移,下一周期会 - # 继续累计;成功切换到新布局后才重新开始统计。本轮异步规划使用独立 - # 的 CPU 快照,不会与推理线程后续的 atomic add 竞争。 + # 此处先不清零 GPU counter:如果样本不足或无需迁移,下一周期会 + # 继续累计;达到重排上限或成功切换到新布局后才重新开始统计。本轮 + # 异步规划使用独立的 CPU 快照,不会与后续的 atomic add 竞争。 local_load = torch.stack([counter.detach().cpu() for counter in counters]) + self._publish_expert_load_metric(local_load) + + # 达到重排次数上限后仍保留周期性负载上报,但不再执行后续的跨 rank + # 通信和布局规划。清空本轮计数,使下一次指标对应新的采样窗口。 + reached_rebalance_limit = ( + self.max_rebalance_count != -1 and self.completed_rebalance_count >= self.max_rebalance_count + ) + if reached_rebalance_limit: + self._clear_route_counters() + self.state = EPLBManagerState.COLLECTING + return # 汇集各 rank 的 token 总数,判断当前统计量是否足以进行布局规划。 token_count_by_rank = [0] * self.world_size @@ -214,7 +221,6 @@ def _step_planning(self) -> None: load_by_rank = list(gathered_load.unbind(dim=0)) dist.all_gather(load_by_rank, local_load, group=self.control_group) global_load = gathered_load.sum(dim=0) - self._publish_expert_load_metric(global_load) self.state = EPLBManagerState.WAIT_PLAN_FINISH if self.global_rank == 0: @@ -293,13 +299,10 @@ def _step_transferring(self) -> None: elapsed = time.time() - self.rebalance_started_at self._clear_route_counters() self.completed_rebalance_count += 1 - reached_rebalance_limit = ( - self.max_rebalance_count != -1 and self.completed_rebalance_count >= self.max_rebalance_count - ) del self.pending_transfer_batches del self.target_placement del self.rebalance_started_at - self.state = EPLBManagerState.FINISHED if reached_rebalance_limit else EPLBManagerState.COLLECTING + self.state = EPLBManagerState.COLLECTING if self.global_rank == 0: logger.info( "eplb completed wall_time=%.2fs completed_rebalance_count=%s max_rebalance_count=%s", @@ -389,12 +392,12 @@ def _publish_layer_metadata(self, layer_index: int) -> None: ) layer_impl.logical_to_physical_map.copy_(logical_to_physical_map) - def _publish_expert_load_metric(self, global_load: torch.Tensor) -> None: + def _publish_expert_load_metric(self, local_load: torch.Tensor) -> None: if self.global_rank != 0: return self.metric_client.gauge_set( EPLB_EXPERT_IMBALANCE_RATIO_METRIC, - _expert_load_imbalance_ratio(global_load), + _expert_load_imbalance_ratio(local_load), ) def _clear_route_counters(self) -> None: @@ -418,14 +421,14 @@ def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) -def _expert_load_imbalance_ratio(global_load: torch.Tensor) -> float: +def _expert_load_imbalance_ratio(expert_load: torch.Tensor) -> float: """计算各层逻辑专家最大 token 数与平均值之比,再对所有层取平均。""" - if global_load.ndim != 2: - raise ValueError("global_load must be [layers, logical_experts]") - global_load = global_load.to(torch.float64) - layer_means = global_load.mean(dim=1) + if expert_load.ndim != 2: + raise ValueError("expert_load must be [layers, logical_experts]") + expert_load = expert_load.to(torch.float64) + layer_means = expert_load.mean(dim=1) valid_layers = layer_means > 0 if not torch.any(valid_layers): return 0.0 - ratios = global_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] + ratios = expert_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] return float(ratios.mean().item()) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index e77066edf4..34a14e39da 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -701,6 +701,8 @@ def test_manager_evaluating_copies_route_counters_to_cpu_without_modifying_them( manager.global_rank = 1 manager.world_size = 1 manager.control_group = object() + manager.max_rebalance_count = -1 + manager.completed_rebalance_count = 0 local_token_counts = [] def all_gather_object(output, local_token_count, **_kwargs): @@ -858,6 +860,8 @@ def test_manager_evaluation_gathers_token_counts_from_all_ranks(monkeypatch): manager.num_redundant_experts_per_rank = 1 manager.current_placement = _initial_expert_placement(4, 4, 1).unsqueeze(0).tolist() manager.control_group = object() + manager.max_rebalance_count = -1 + manager.completed_rebalance_count = 0 local = torch.full((4,), 100, dtype=torch.int64) manager._eplb_impls[0].route_counter = local manager.state = manager_module.EPLBManagerState.EVALUATING @@ -1482,7 +1486,7 @@ def all_gather_object(output, local_state, **_kwargs): assert waits == [overlap_stream, overlap_stream] -def test_manager_finishes_after_reaching_rebalance_limit(): +def test_manager_returns_to_collecting_after_reaching_rebalance_limit(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) impls = [SimpleNamespace(recording=True), SimpleNamespace(recording=True)] target_placement = [[[0, 1], [1, 0]]] @@ -1500,7 +1504,7 @@ def test_manager_finishes_after_reaching_rebalance_limit(): assert manager.current_placement is target_placement assert manager.completed_rebalance_count == 1 - assert manager.state is manager_module.EPLBManagerState.FINISHED + assert manager.state is manager_module.EPLBManagerState.COLLECTING assert all(impl.recording for impl in impls) @@ -1603,15 +1607,6 @@ def test_manager_step_advances_inflight_transfer(): assert calls == ["transfer"] -def test_finished_manager_step_is_a_noop(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.state = manager_module.EPLBManagerState.FINISHED - - manager.step() - - assert manager.state is manager_module.EPLBManagerState.FINISHED - - def test_manager_evaluates_only_after_entering_evaluating_state(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) route_counter = torch.tensor([1, 2], dtype=torch.int64) @@ -1625,6 +1620,8 @@ def test_manager_evaluates_only_after_entering_evaluating_state(monkeypatch): manager._eplb_impls = [SimpleNamespace(route_counter=route_counter)] manager.world_size = 1 manager.control_group = object() + manager.max_rebalance_count = -1 + manager.completed_rebalance_count = 0 def all_gather_object(output, local_token_count, **_kwargs): local_token_counts.append(local_token_count) @@ -1725,6 +1722,8 @@ def test_manager_evaluation_with_insufficient_tokens_returns_to_collecting(monke manager.global_rank = 1 manager.world_size = 1 manager.control_group = object() + manager.max_rebalance_count = -1 + manager.completed_rebalance_count = 0 monkeypatch.setattr( manager_module.dist, "all_gather_object", @@ -1741,12 +1740,15 @@ def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) local_load = torch.full((1, 4), 256, dtype=torch.int64) plan_tasks = [] + published_loads = [] manager.state = manager_module.EPLBManagerState.EVALUATING manager.global_rank = 0 manager.num_logical_experts = 4 manager._eplb_impls = [SimpleNamespace(route_counter=local_load[0])] manager.world_size = 1 manager.control_group = object() + manager.max_rebalance_count = -1 + manager.completed_rebalance_count = 0 monkeypatch.setattr( manager_module.dist, "all_gather_object", @@ -1771,13 +1773,15 @@ def start(self): manager.planner = object() manager.current_placement = [[[0, 1, 2, 3]]] - manager._publish_expert_load_metric = lambda _global_load: None + manager._publish_expert_load_metric = lambda load: published_loads.append(load) monkeypatch.setattr(manager_module, "EPLBPlanTask", PlanTask) manager.step() assert manager.state is manager_module.EPLBManagerState.PLANNING assert torch.equal(manager._local_load, local_load) + assert len(published_loads) == 1 + assert torch.equal(published_loads[0], local_load) assert plan_tasks == [] assert not hasattr(manager, "_plan_task") @@ -1790,6 +1794,34 @@ def start(self): assert plan_tasks[0].started assert manager._plan_task is plan_tasks[0] assert not hasattr(manager, "_local_load") + assert len(published_loads) == 1 + + +def test_manager_keeps_reporting_after_reaching_rebalance_limit(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + local_load = torch.tensor([[10, 20, 30, 40]], dtype=torch.int64) + published_loads = [] + cleared_counters = [] + manager.state = manager_module.EPLBManagerState.EVALUATING + manager.global_rank = 0 + manager.num_logical_experts = 4 + manager._eplb_impls = [SimpleNamespace(route_counter=local_load[0])] + manager.max_rebalance_count = 1 + manager.completed_rebalance_count = 1 + manager._publish_expert_load_metric = lambda load: published_loads.append(load) + manager._clear_route_counters = lambda: cleared_counters.append(True) + monkeypatch.setattr( + manager_module.dist, + "all_gather_object", + lambda *_args, **_kwargs: pytest.fail("rebalancing must stop after reaching the limit"), + ) + + manager.step() + + assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert len(published_loads) == 1 + assert torch.equal(published_loads[0], local_load) + assert cleared_counters == [True] def test_nonzero_rank_waits_without_starting_planner(monkeypatch): From 60df60898bd15b320c4fd91f971bfe798bcf2445 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 02:27:03 +0000 Subject: [PATCH 46/72] clarify EPLB manager control flow --- .../model_infer/mode_backend/eplb_manager.py | 111 +++++++++--------- 1 file changed, 54 insertions(+), 57 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index ca60ec4d80..558f388455 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -158,9 +158,9 @@ def _step_collecting(self) -> None: self.steps += 1 if self.steps < self.next_evaluation_step: return - - self.next_evaluation_step += self.step_interval - self.state = EPLBManagerState.EVALUATING + else: + self.next_evaluation_step += self.step_interval + self.state = EPLBManagerState.EVALUATING def _step_evaluating(self) -> None: """发布本地负载指标,并在次数允许时根据全局样本量决定是否规划。""" @@ -183,28 +183,26 @@ def _step_evaluating(self) -> None: if reached_rebalance_limit: self._clear_route_counters() self.state = EPLBManagerState.COLLECTING - return - - # 汇集各 rank 的 token 总数,判断当前统计量是否足以进行布局规划。 - token_count_by_rank = [0] * self.world_size - dist.all_gather_object( - token_count_by_rank, - int(local_load.sum().item()), - group=self.control_group, - ) - average_tokens_per_expert = sum(token_count_by_rank) / local_load.numel() - if average_tokens_per_expert < EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT: - if self.global_rank == 0: - logger.info( - "eplb continue collecting average_tokens_per_expert=%.2f threshold=%s", - average_tokens_per_expert, - EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT, - ) - self.state = EPLBManagerState.COLLECTING - return - - self._local_load = local_load - self.state = EPLBManagerState.PLANNING + else: + # 汇集各 rank 的 token 总数,判断当前统计量是否足以进行布局规划。 + token_count_by_rank = [0] * self.world_size + dist.all_gather_object( + token_count_by_rank, + int(local_load.sum().item()), + group=self.control_group, + ) + average_tokens_per_expert = sum(token_count_by_rank) / local_load.numel() + if average_tokens_per_expert < EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT: + if self.global_rank == 0: + logger.info( + "eplb continue collecting average_tokens_per_expert=%.2f threshold=%s", + average_tokens_per_expert, + EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT, + ) + self.state = EPLBManagerState.COLLECTING + else: + self._local_load = local_load + self.state = EPLBManagerState.PLANNING def _step_planning(self) -> None: """汇集全局负载,并由 rank 0 启动异步规划。""" @@ -310,25 +308,24 @@ def _step_transferring(self) -> None: self.completed_rebalance_count, self.max_rebalance_count, ) - return - - self.active_transfer_batch = transfer_batch - - # 普通批次允许多个 rank 不冲突的任务并行,但每个 rank 最多参与 - # 一条;覆盖环批次可能要求同一 rank 同时保存多个源/目标的 pinned - # row,必须等整批传输完成后再统一覆盖 live 权重。 - self.active_transfers = [ - PinnedMemoryEPLBTransfer( - self._weights, - self.transfer_group, - self.global_rank, - transfer_info, - ) - for transfer_info in transfer_batch - if self.global_rank in (transfer_info.source_rank, transfer_info.dest_rank) - ] - for transfer in self.active_transfers: - transfer.start() + else: + self.active_transfer_batch = transfer_batch + + # 普通批次允许多个 rank 不冲突的任务并行,但每个 rank 最多参与 + # 一条;覆盖环批次可能要求同一 rank 同时保存多个源/目标的 pinned + # row,必须等整批传输完成后再统一覆盖 live 权重。 + self.active_transfers = [ + PinnedMemoryEPLBTransfer( + self._weights, + self.transfer_group, + self.global_rank, + transfer_info, + ) + for transfer_info in transfer_batch + if self.global_rank in (transfer_info.source_rank, transfer_info.dest_rank) + ] + for transfer in self.active_transfers: + transfer.start() else: # 已有活动批次时,本 step 只负责轮询;整批完成后才统一提交。 self._poll_transfer_batch() @@ -343,19 +340,19 @@ def _poll_transfer_batch(self) -> None: dist.all_gather_object(finished_by_rank, local_finished, group=self.control_group) if not all(finished_by_rank): return - - # 所有 rank 使用相同的批次顺序提交,因此全局 placement 和 metadata - # 始终一致;只有 destination rank 会额外写入实际专家权重。 - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - for transfer_info in self.active_transfer_batch: - self._commit_transfer(transfer_info) - for layer_index in {transfer_info.layer_index for transfer_info in self.active_transfer_batch}: - self._publish_layer_metadata(layer_index) - - del self.active_transfers - del self.active_transfer_batch + else: + # 所有 rank 使用相同的批次顺序提交,因此全局 placement 和 metadata + # 始终一致;只有 destination rank 会额外写入实际专家权重。 + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) + for transfer_info in self.active_transfer_batch: + self._commit_transfer(transfer_info) + for layer_index in {transfer_info.layer_index for transfer_info in self.active_transfer_batch}: + self._publish_layer_metadata(layer_index) + + del self.active_transfers + del self.active_transfer_batch def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: """把一条已完成传输提交到 live 权重和完整布局。""" From a452432a369e3f19ee9a7fb03df9f4df1cb42e5b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 03:01:17 +0000 Subject: [PATCH 47/72] plan EPLB transfers asynchronously --- .../model_infer/mode_backend/eplb_manager.py | 112 +++++++++----- .../mode_backend/eplb_transfer_planner.py | 77 ++++++++++ unit_tests/common/fused_moe/test_eplb.py | 141 +++++++++++++++--- 3 files changed, 267 insertions(+), 63 deletions(-) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb_transfer_planner.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index 558f388455..8f53c67253 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -22,8 +22,8 @@ from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( EPLBTransferInfo, PinnedMemoryEPLBTransfer, - build_transfer_plan, ) +from lightllm.server.router.model_infer.mode_backend.eplb_transfer_planner import EPLBTransferPlanner from lightllm.utils.dist_utils import ( get_global_rank, get_global_world_size, @@ -43,8 +43,10 @@ class EPLBManagerState(Enum): COLLECTING = "collecting" EVALUATING = "evaluating" - PLANNING = "planning" - WAIT_PLAN_FINISH = "wait_plan_finish" + PLAN_PLACEMENT = "plan_placement" + WAIT_PLAN_PLACEMENT_FINISHED = "wait_plan_placement_finished" + PLAN_TRANSFER = "plan_transfer" + WAIT_PLAN_TRANSFER_FINISHED = "wait_plan_transfer_finished" TRANSFERRING = "transferring" @@ -53,14 +55,16 @@ class EPLBManager: 状态循环如下: - ``COLLECTING -> EVALUATING -> PLANNING -> WAIT_PLAN_FINISH -> TRANSFERRING`` + ``COLLECTING -> EVALUATING -> PLAN_PLACEMENT -> WAIT_PLAN_PLACEMENT_FINISHED`` + ``-> PLAN_TRANSFER -> WAIT_PLAN_TRANSFER_FINISHED -> TRANSFERRING`` 当累计的平均专家 token 数不足时,``EVALUATING`` 会回到 - ``COLLECTING``;当规划器认为无需调整布局时,``WAIT_PLAN_FINISH`` 会 + ``COLLECTING``;当规划器认为无需调整布局时, + ``WAIT_PLAN_PLACEMENT_FINISHED`` 会 回到 ``COLLECTING``。达到重排次数上限后,manager 仍会周期性采集并 上报本地负载指标,但不再进入全局负载汇总和布局规划。每次调用 - :meth:`step` 最多推进一个状态,布局规划和权重传输在后台执行,主推理 - 线程负责评估、轮询和提交结果。 + :meth:`step` 最多推进一个状态,布局规划、传输规划和权重传输都在后台 + 执行,主推理线程负责评估、轮询和提交结果。 """ def __init__(self, model: TpPartBaseModel, max_rebalance_count: int = 1) -> None: @@ -137,12 +141,20 @@ def step(self) -> None: self._step_evaluating() return - if self.state is EPLBManagerState.PLANNING: - self._step_planning() + if self.state is EPLBManagerState.PLAN_PLACEMENT: + self._step_plan_placement() return - if self.state is EPLBManagerState.WAIT_PLAN_FINISH: - self._step_wait_plan_finish() + if self.state is EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED: + self._step_wait_plan_placement_finished() + return + + if self.state is EPLBManagerState.PLAN_TRANSFER: + self._step_plan_transfer() + return + + if self.state is EPLBManagerState.WAIT_PLAN_TRANSFER_FINISHED: + self._step_wait_plan_transfer_finished() return if self.state is EPLBManagerState.TRANSFERRING: @@ -202,9 +214,9 @@ def _step_evaluating(self) -> None: self.state = EPLBManagerState.COLLECTING else: self._local_load = local_load - self.state = EPLBManagerState.PLANNING + self.state = EPLBManagerState.PLAN_PLACEMENT - def _step_planning(self) -> None: + def _step_plan_placement(self) -> None: """汇集全局负载,并由 rank 0 启动异步规划。""" local_load = self._local_load del self._local_load @@ -220,7 +232,7 @@ def _step_planning(self) -> None: dist.all_gather(load_by_rank, local_load, group=self.control_group) global_load = gathered_load.sum(dim=0) - self.state = EPLBManagerState.WAIT_PLAN_FINISH + self.state = EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED if self.global_rank == 0: self._plan_task = EPLBPlanTask( self.planner, @@ -229,7 +241,7 @@ def _step_planning(self) -> None: ) self._plan_task.start() - def _step_wait_plan_finish(self) -> None: + def _step_wait_plan_placement_finished(self) -> None: """等待 rank 0 完成规划并广播目标专家排布。""" placement: Optional[ExpertPlacement] = None if self.global_rank == 0 and self._plan_task.is_finished(): @@ -254,33 +266,51 @@ def _step_wait_plan_finish(self) -> None: # 广播得到的 placement 已经是规划器新建的完整布局,没有外部持有者会 # 再修改它,因此可直接保存,不需要逐层深拷贝。 self.target_placement = placement + self.state = EPLBManagerState.PLAN_TRANSFER + + def _step_plan_transfer(self) -> None: + """启动异步传输规划,再进入完成状态轮询阶段。""" + # 所有 rank 使用相同的 current/target placement 独立生成确定性的传输 + # 批次,避免广播体积较大的任务列表;耗时的逐层依赖分析放到后台线程, + # 当前推理线程从下一次安全边界开始只需轮询完成状态。 + self._transfer_planner = EPLBTransferPlanner( + self.current_placement, + self.target_placement, + self.num_logical_experts, + self.world_size, + ) + self._transfer_planner.start() + self.state = EPLBManagerState.WAIT_PLAN_TRANSFER_FINISHED - # 每层独立构建有序传输批次,再按 layer 顺序拼接。这样一个批次内只 - # 包含同层任务,提交完成后也只需发布该层的路由 metadata。 - self.pending_transfer_batches: List[List[EPLBTransferInfo]] = [] - layer_placements = zip(self.current_placement, self.target_placement) - for layer_index, (current_layer, target_layer) in enumerate(layer_placements): - layer_transfer_batches = build_transfer_plan( - current_layer, - target_layer, - layer_index, - self.num_logical_experts, - self.world_size, - ) - self.pending_transfer_batches.extend(layer_transfer_batches) - if not self.pending_transfer_batches: - raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") - self.state = EPLBManagerState.TRANSFERRING - if self.global_rank == 0: - changed_layer_count = sum( - current != target for current, target in zip(self.current_placement, self.target_placement) - ) - logger.info( - "eplb started steps=%s changed_layer_count=%s changed_slot_count=%s", - self.steps, - changed_layer_count, - sum(len(transfer_batch) for transfer_batch in self.pending_transfer_batches), - ) + def _step_wait_plan_transfer_finished(self) -> None: + """等待所有 rank 异步生成相同的传输批次,再统一进入传输状态。""" + local_finished = self._transfer_planner.is_finished() + finished_by_rank = [False] * self.world_size + dist.all_gather_object(finished_by_rank, local_finished, group=self.control_group) + + # 即使本 rank 已经完成,也必须等待其他 rank 的镜像任务列表就绪;否则 + # 提前进入 TRANSFERRING 的 rank 可能发起尚无对端参与的点对点传输。 + if not all(finished_by_rank): + return + else: + pending_transfer_batches = self._transfer_planner.result + assert pending_transfer_batches is not None + self.pending_transfer_batches = pending_transfer_batches + del self._transfer_planner + if not self.pending_transfer_batches: + raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") + + self.state = EPLBManagerState.TRANSFERRING + if self.global_rank == 0: + changed_layer_count = sum( + current != target for current, target in zip(self.current_placement, self.target_placement) + ) + logger.info( + "eplb started steps=%s changed_layer_count=%s changed_slot_count=%s", + self.steps, + changed_layer_count, + sum(len(transfer_batch) for transfer_batch in self.pending_transfer_batches), + ) def _step_transferring(self) -> None: """启动或轮询一个传输批次;整批完成后再统一提交。""" diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer_planner.py new file mode 100644 index 0000000000..f78df74669 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer_planner.py @@ -0,0 +1,77 @@ +"""EPLB 专家传输计划的异步生成器。""" + +import os +import threading +from enum import Enum +from typing import List, Optional + +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ExpertPlacement +from lightllm.server.router.model_infer.mode_backend.eplb_transfer import EPLBTransferInfo, build_transfer_plan +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +class TransferPlanStatus(Enum): + """异步传输规划器的生命周期状态。""" + + IDLE = "idle" + RUNNING = "running" + SUCCEEDED = "succeeded" + + +class EPLBTransferPlanner: + """在后台线程中逐层生成并按 layer 顺序拼接专家传输批次。 + + 输入布局在规划期间保持只读。每层独立调用 ``build_transfer_plan``,因此 + 一个批次只包含同一层的任务,manager 可以在提交后立即发布该层的路由 + metadata。 + """ + + def __init__( + self, + current_placement: ExpertPlacement, + target_placement: ExpertPlacement, + num_logical_experts: int, + world_size: int, + ) -> None: + self.current_placement = current_placement + self.target_placement = target_placement + self.num_logical_experts = num_logical_experts + self.world_size = world_size + self.status = TransferPlanStatus.IDLE + self.result: Optional[List[List[EPLBTransferInfo]]] = None + self._thread = threading.Thread( + target=self._run, + name="eplb-transfer-plan", + daemon=True, + ) + + def start(self) -> None: + """启动异步传输规划。""" + assert self.status is TransferPlanStatus.IDLE, "EPLB transfer planner has already been started" + self.status = TransferPlanStatus.RUNNING + self._thread.start() + + def is_finished(self) -> bool: + """返回全部层的传输批次是否已经生成。""" + return self.status is TransferPlanStatus.SUCCEEDED + + def _run(self) -> None: + try: + transfer_batches: List[List[EPLBTransferInfo]] = [] + layer_placements = zip(self.current_placement, self.target_placement) + for layer_index, (current_layer, target_layer) in enumerate(layer_placements): + layer_transfer_batches = build_transfer_plan( + current_layer, + target_layer, + layer_index, + self.num_logical_experts, + self.world_size, + ) + transfer_batches.extend(layer_transfer_batches) + self.result = transfer_batches + self.status = TransferPlanStatus.SUCCEEDED + except BaseException: + logger.exception("EPLB transfer planning failed") + os._exit(1) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 34a14e39da..adf67baff2 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -27,6 +27,9 @@ from lightllm.server.router.model_infer.mode_backend import ( eplb_transfer as transfer_module, ) +from lightllm.server.router.model_infer.mode_backend import ( + eplb_transfer_planner as transfer_planner_module, +) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( deepgemm_impl as deepgemm_module, ) @@ -755,6 +758,59 @@ def fail(_load, _placement): assert logs == ["EPLB planning failed"] +def test_transfer_planner_combines_all_layer_batches(monkeypatch): + current_placement = [[[0, 1], [2, 3]], [[0, 2], [1, 3]]] + target_placement = [[[2, 1], [0, 3]], [[0, 3], [1, 2]]] + transfer_infos = [ + EPLBTransferInfo(1, 0, 2, 0, 0), + EPLBTransferInfo(1, 1, 3, 0, 1), + ] + calls = [] + + def build_plan(*args): + calls.append(args) + return [[transfer_infos[args[2]]]] + + monkeypatch.setattr(transfer_planner_module, "build_transfer_plan", build_plan) + planner = transfer_planner_module.EPLBTransferPlanner( + current_placement, + target_placement, + num_logical_experts=4, + world_size=2, + ) + + planner._run() + + assert planner.status is transfer_planner_module.TransferPlanStatus.SUCCEEDED + assert planner.result == [[transfer_infos[0]], [transfer_infos[1]]] + assert calls == [ + (current_placement[0], target_placement[0], 0, 4, 2), + (current_placement[1], target_placement[1], 1, 4, 2), + ] + + +def test_transfer_planner_exits_process_on_failure(monkeypatch): + def fail(*_args): + raise RuntimeError("transfer planning boom") + + planner = transfer_planner_module.EPLBTransferPlanner( + [[[0], [1]]], + [[[1], [0]]], + num_logical_experts=2, + world_size=2, + ) + exits = [] + logs = [] + monkeypatch.setattr(transfer_planner_module, "build_transfer_plan", fail) + monkeypatch.setattr(transfer_planner_module.os, "_exit", exits.append) + monkeypatch.setattr(transfer_planner_module.logger, "exception", logs.append) + + planner._run() + + assert exits == [1] + assert logs == ["EPLB transfer planning failed"] + + def test_expert_load_imbalance_ratio_averages_layer_ratios(): global_load = torch.tensor( [ @@ -1510,7 +1566,7 @@ def test_manager_returns_to_collecting_after_reaching_rebalance_limit(): def test_wait_plan_finish_broadcasts_pending_status(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.state = manager_module.EPLBManagerState.WAIT_PLAN_FINISH + manager.state = manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED manager.global_rank = 0 manager.control_group = object() manager._plan_task = SimpleNamespace( @@ -1523,9 +1579,9 @@ def broadcast(values, **_kwargs): broadcasts.append(values[0]) monkeypatch.setattr(manager_module.dist, "broadcast_object_list", broadcast) - manager._step_wait_plan_finish() + manager._step_wait_plan_placement_finished() assert broadcasts == [None] - assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @@ -1660,9 +1716,9 @@ def test_manager_step_uses_explicit_state_instead_of_pending_work(): assert calls == ["transfer"] -def test_manager_enters_transferring_state_with_planned_work(monkeypatch): +def test_manager_plans_transfers_asynchronously_before_entering_transferring(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.state = manager_module.EPLBManagerState.WAIT_PLAN_FINISH + manager.state = manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED manager.global_rank = 1 manager.control_group = object() manager.transfer_group = object() @@ -1686,27 +1742,68 @@ def test_manager_enters_transferring_state_with_planned_work(monkeypatch): EPLBTransferInfo(0, 0, 2, 1, 2), EPLBTransferInfo(0, 1, 0, 1, 2), ] - build_calls = [] + transfer_planners = [] + + class TransferPlanner: + def __init__(self, current, target, num_logical_experts, world_size): + self.current = current + self.target = target + self.num_logical_experts = num_logical_experts + self.world_size = world_size + self.result = [[transfer_infos[0]], [transfer_infos[1]]] + self.started = False + self.finished = False + transfer_planners.append(self) - def build_plan(*args): - build_calls.append(args) - return [[transfer_infos[args[2]]]] + def start(self): + self.started = True + + def is_finished(self): + return self.finished - monkeypatch.setattr(manager_module, "build_transfer_plan", build_plan) + monkeypatch.setattr(manager_module, "EPLBTransferPlanner", TransferPlanner) monkeypatch.setattr( manager_module, "PinnedMemoryEPLBTransfer", - lambda *_args: pytest.fail("transfer object must not be built while entering TRANSFERRING"), + lambda *_args: pytest.fail("transfer object must not be built while planning transfers"), ) - manager._step_wait_plan_finish() + manager._step_wait_plan_placement_finished() + + assert manager.state is manager_module.EPLBManagerState.PLAN_TRANSFER + assert manager.target_placement is placement + assert transfer_planners == [] + assert not hasattr(manager, "pending_transfer_batches") + + manager.step() + + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_TRANSFER_FINISHED + assert len(transfer_planners) == 1 + transfer_planner = transfer_planners[0] + assert transfer_planner.current is manager.current_placement + assert transfer_planner.target is placement + assert transfer_planner.num_logical_experts == manager.num_logical_experts + assert transfer_planner.world_size == manager.world_size + assert transfer_planner.started + + remote_finished = False + + def gather_finished(output, local_finished, **_kwargs): + output[:] = [local_finished, remote_finished] + + monkeypatch.setattr(manager_module.dist, "all_gather_object", gather_finished) + transfer_planner.finished = True + manager.step() + + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_TRANSFER_FINISHED + assert not hasattr(manager, "pending_transfer_batches") + + remote_finished = True + manager.step() assert manager.state is manager_module.EPLBManagerState.TRANSFERRING assert manager.pending_transfer_batches == [[transfer_infos[0]], [transfer_infos[1]]] - assert build_calls == [ - (manager.current_placement[0], placement[0], 0, manager.num_logical_experts, manager.world_size), - (manager.current_placement[1], placement[1], 1, manager.num_logical_experts, manager.world_size), - ] + assert not hasattr(manager, "_transfer_planner") assert not hasattr(manager, "active_transfer") assert not hasattr(manager, "target_metadata") @@ -1778,7 +1875,7 @@ def start(self): manager.step() - assert manager.state is manager_module.EPLBManagerState.PLANNING + assert manager.state is manager_module.EPLBManagerState.PLAN_PLACEMENT assert torch.equal(manager._local_load, local_load) assert len(published_loads) == 1 assert torch.equal(published_loads[0], local_load) @@ -1787,7 +1884,7 @@ def start(self): manager.step() - assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED assert torch.equal(plan_tasks[0].global_load, local_load) assert plan_tasks[0].planner is manager.planner assert plan_tasks[0].current_placement is manager.current_placement @@ -1826,7 +1923,7 @@ def test_manager_keeps_reporting_after_reaching_rebalance_limit(monkeypatch): def test_nonzero_rank_waits_without_starting_planner(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.state = manager_module.EPLBManagerState.PLANNING + manager.state = manager_module.EPLBManagerState.PLAN_PLACEMENT manager.global_rank = 1 manager.world_size = 2 manager.control_group = object() @@ -1840,14 +1937,14 @@ def all_gather(output, local, **_kwargs): manager.step() - assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED assert not hasattr(manager, "_plan_task") assert not hasattr(manager, "_local_load") def test_manager_planning_without_changes_returns_to_collecting(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.state = manager_module.EPLBManagerState.WAIT_PLAN_FINISH + manager.state = manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED manager.global_rank = 1 manager.control_group = object() manager.steps = 11 @@ -1862,7 +1959,7 @@ def broadcast(values, **_kwargs): monkeypatch.setattr(manager_module.dist, "broadcast_object_list", broadcast) manager.step() - assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_FINISH + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED result = [[[0, 1], [1, 0]]] manager.step() From b85aa7474ce300c155a473f55c93e16c59c7339f Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 03:05:20 +0000 Subject: [PATCH 48/72] group EPLB runtime modules into package --- .../model_infer/mode_backend/base_backend.py | 2 +- .../model_infer/mode_backend/eplb/__init__.py | 1 + .../{eplb_manager.py => eplb/manager.py} | 6 +++--- .../{eplb_plan.py => eplb/plan.py} | 0 .../{eplb_transfer.py => eplb/transfer.py} | 0 .../transfer_planner.py} | 2 +- unit_tests/common/fused_moe/test_eplb.py | 18 +++++++++--------- .../common/fused_moe/test_eplb_transfer_gpu.py | 2 +- 8 files changed, 16 insertions(+), 15 deletions(-) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/__init__.py rename lightllm/server/router/model_infer/mode_backend/{eplb_manager.py => eplb/manager.py} (98%) rename lightllm/server/router/model_infer/mode_backend/{eplb_plan.py => eplb/plan.py} (100%) rename lightllm/server/router/model_infer/mode_backend/{eplb_transfer.py => eplb/transfer.py} (100%) rename lightllm/server/router/model_infer/mode_backend/{eplb_transfer_planner.py => eplb/transfer_planner.py} (96%) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index ec85ef535f..8f7b642d8b 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -258,7 +258,7 @@ def init_model(self, kvargs): prof_mode = self.args.enable_profiling self.profiler = ProcessProfiler(mode=prof_mode, name=prof_name, use_multi_thread=True) if prof_mode else None if self.args.eplb_num_redundant_experts_per_rank > 0: - from lightllm.server.router.model_infer.mode_backend.eplb_manager import EPLBManager + from lightllm.server.router.model_infer.mode_backend.eplb.manager import EPLBManager self.eplb_manager = EPLBManager(self.model, self.args.eplb_rebalance_count) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/__init__.py b/lightllm/server/router/model_infer/mode_backend/eplb/__init__.py new file mode 100644 index 0000000000..ca9366913f --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/__init__.py @@ -0,0 +1 @@ +"""EPLB 布局规划、传输规划和运行时管理。""" diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/manager.py similarity index 98% rename from lightllm/server/router/model_infer/mode_backend/eplb_manager.py rename to lightllm/server/router/model_infer/mode_backend/eplb/manager.py index 8f53c67253..e1e072d0c0 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/manager.py @@ -18,12 +18,12 @@ FusedMoeWeight, ) from lightllm.server.metrics.manager import MetricClient -from lightllm.server.router.model_infer.mode_backend.eplb_plan import EPLBPlanTask -from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( +from .plan import EPLBPlanTask +from .transfer import ( EPLBTransferInfo, PinnedMemoryEPLBTransfer, ) -from lightllm.server.router.model_infer.mode_backend.eplb_transfer_planner import EPLBTransferPlanner +from .transfer_planner import EPLBTransferPlanner from lightllm.utils.dist_utils import ( get_global_rank, get_global_world_size, diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_plan.py b/lightllm/server/router/model_infer/mode_backend/eplb/plan.py similarity index 100% rename from lightllm/server/router/model_infer/mode_backend/eplb_plan.py rename to lightllm/server/router/model_infer/mode_backend/eplb/plan.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb/transfer.py similarity index 100% rename from lightllm/server/router/model_infer/mode_backend/eplb_transfer.py rename to lightllm/server/router/model_infer/mode_backend/eplb/transfer.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py similarity index 96% rename from lightllm/server/router/model_infer/mode_backend/eplb_transfer_planner.py rename to lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py index f78df74669..1b973f7259 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer_planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py @@ -6,7 +6,7 @@ from typing import List, Optional from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ExpertPlacement -from lightllm.server.router.model_infer.mode_backend.eplb_transfer import EPLBTransferInfo, build_transfer_plan +from .transfer import EPLBTransferInfo, build_transfer_plan from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index adf67baff2..77b15fbbe1 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -18,17 +18,17 @@ from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.server.router.model_infer.infer_batch import g_infer_context -from lightllm.server.router.model_infer.mode_backend import ( - eplb_manager as manager_module, +from lightllm.server.router.model_infer.mode_backend.eplb import ( + manager as manager_module, ) -from lightllm.server.router.model_infer.mode_backend import ( - eplb_plan as plan_module, +from lightllm.server.router.model_infer.mode_backend.eplb import ( + plan as plan_module, ) -from lightllm.server.router.model_infer.mode_backend import ( - eplb_transfer as transfer_module, +from lightllm.server.router.model_infer.mode_backend.eplb import ( + transfer as transfer_module, ) -from lightllm.server.router.model_infer.mode_backend import ( - eplb_transfer_planner as transfer_planner_module, +from lightllm.server.router.model_infer.mode_backend.eplb import ( + transfer_planner as transfer_planner_module, ) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( deepgemm_impl as deepgemm_module, @@ -45,7 +45,7 @@ fused_moe_weight as fused_weight_module, ) from lightllm.common.eplb_utils import extract_eplb_expert_tensors -from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( +from lightllm.server.router.model_infer.mode_backend.eplb.transfer import ( EPLBTransferInfo, ExpertTensorBuffer, PinnedMemoryEPLBTransfer, diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index c0ae0b856f..17e63b3eb6 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -11,7 +11,7 @@ import torch.distributed as dist import torch.multiprocessing as mp -from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( +from lightllm.server.router.model_infer.mode_backend.eplb.transfer import ( PinnedMemoryEPLBTransfer, TransferStatus, build_transfer_plan, From 20f924ac8d2063ab42b077ee30e3a8921b01e603 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 03:11:53 +0000 Subject: [PATCH 49/72] move EPLB placement and planner into runtime package --- .../fused_moe/impl/deepgemm_impl.py | 2 +- .../model_infer/mode_backend/eplb/manager.py | 23 ++++++++----------- .../mode_backend/eplb/placement.py} | 0 .../model_infer/mode_backend/eplb/plan.py | 6 ++--- .../model_infer/mode_backend/eplb/planner.py} | 0 .../mode_backend/eplb/transfer_planner.py | 5 ++-- unit_tests/common/fused_moe/test_eplb.py | 4 ++-- 7 files changed, 17 insertions(+), 23 deletions(-) rename lightllm/{common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py => server/router/model_infer/mode_backend/eplb/placement.py} (100%) rename lightllm/{common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py => server/router/model_infer/mode_backend/eplb/planner.py} (100%) 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 92ba5c4677..0c9a6e3b3a 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 @@ -1,7 +1,7 @@ import torch from typing import Optional, Tuple, Any from .base_impl import FuseMoeBaseImpl -from ..eplb_placement import ( +from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( build_initial_local_expert_ids, build_logical_to_physical_map, ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/manager.py index e1e072d0c0..611eed1810 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/manager.py @@ -6,24 +6,10 @@ import torch.distributed as dist from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_logical_to_physical_map, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ( - EPLBPlanner, - ExpertPlacement, - GreedyEPLBPlanner, -) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import ( FusedMoeWeight, ) from lightllm.server.metrics.manager import MetricClient -from .plan import EPLBPlanTask -from .transfer import ( - EPLBTransferInfo, - PinnedMemoryEPLBTransfer, -) -from .transfer_planner import EPLBTransferPlanner from lightllm.utils.dist_utils import ( get_global_rank, get_global_world_size, @@ -32,6 +18,15 @@ from lightllm.utils.log_utils import init_logger from lightllm.utils.shm_port_args import get_shm_port_args +from .placement import build_logical_to_physical_map +from .plan import EPLBPlanTask +from .planner import EPLBPlanner, ExpertPlacement, GreedyEPLBPlanner +from .transfer import ( + EPLBTransferInfo, + PinnedMemoryEPLBTransfer, +) +from .transfer_planner import EPLBTransferPlanner + logger = init_logger(__name__) EPLB_EXPERT_ALIGNMENT = 128 EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT = 256 diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement.py similarity index 100% rename from lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py rename to lightllm/server/router/model_infer/mode_backend/eplb/placement.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/plan.py b/lightllm/server/router/model_infer/mode_backend/eplb/plan.py index b6400a9122..216c4cba59 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/plan.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/plan.py @@ -7,12 +7,10 @@ import torch -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ( - EPLBPlanner, - ExpertPlacement, -) from lightllm.utils.log_utils import init_logger +from .planner import EPLBPlanner, ExpertPlacement + logger = init_logger(__name__) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/planner.py similarity index 100% rename from lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_planner.py rename to lightllm/server/router/model_infer/mode_backend/eplb/planner.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py index 1b973f7259..6102158ebf 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py @@ -5,10 +5,11 @@ from enum import Enum from typing import List, Optional -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ExpertPlacement -from .transfer import EPLBTransferInfo, build_transfer_plan from lightllm.utils.log_utils import init_logger +from .planner import ExpertPlacement +from .transfer import EPLBTransferInfo, build_transfer_plan + logger = init_logger(__name__) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 77b15fbbe1..bf60ddf7ba 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -6,12 +6,12 @@ import pytest import torch -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( +from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( build_initial_local_expert_ids, build_logical_to_physical_map, build_logical_to_physical_maps_for_layers, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_planner import ( +from lightllm.server.router.model_infer.mode_backend.eplb.planner import ( EPLBPlanner, GreedyEPLBPlanner, ) From 691d44596377c8f6ad3df2ba26b0392eef1399ad Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 03:16:53 +0000 Subject: [PATCH 50/72] clarify EPLB runtime module names --- .../meta_weights/fused_moe/impl/deepgemm_impl.py | 2 +- .../model_infer/mode_backend/base_backend.py | 2 +- ...ansfer_planner.py => async_transfer_planner.py} | 4 ++-- .../eplb/{placement.py => expert_placement.py} | 0 .../eplb/{transfer.py => expert_transfer.py} | 0 .../eplb/{plan.py => placement_plan_task.py} | 2 +- .../eplb/{planner.py => placement_planner.py} | 0 .../eplb/{manager.py => runtime_manager.py} | 10 +++++----- unit_tests/common/fused_moe/test_eplb.py | 14 +++++++------- .../common/fused_moe/test_eplb_transfer_gpu.py | 2 +- 10 files changed, 18 insertions(+), 18 deletions(-) rename lightllm/server/router/model_infer/mode_backend/eplb/{transfer_planner.py => async_transfer_planner.py} (95%) rename lightllm/server/router/model_infer/mode_backend/eplb/{placement.py => expert_placement.py} (100%) rename lightllm/server/router/model_infer/mode_backend/eplb/{transfer.py => expert_transfer.py} (100%) rename lightllm/server/router/model_infer/mode_backend/eplb/{plan.py => placement_plan_task.py} (96%) rename lightllm/server/router/model_infer/mode_backend/eplb/{planner.py => placement_planner.py} (100%) rename lightllm/server/router/model_infer/mode_backend/eplb/{manager.py => runtime_manager.py} (98%) 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 0c9a6e3b3a..7da12de0a0 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 @@ -1,7 +1,7 @@ import torch from typing import Optional, Tuple, Any from .base_impl import FuseMoeBaseImpl -from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( +from lightllm.server.router.model_infer.mode_backend.eplb.expert_placement import ( build_initial_local_expert_ids, build_logical_to_physical_map, ) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 8f7b642d8b..afbbc1ca49 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -258,7 +258,7 @@ def init_model(self, kvargs): prof_mode = self.args.enable_profiling self.profiler = ProcessProfiler(mode=prof_mode, name=prof_name, use_multi_thread=True) if prof_mode else None if self.args.eplb_num_redundant_experts_per_rank > 0: - from lightllm.server.router.model_infer.mode_backend.eplb.manager import EPLBManager + from lightllm.server.router.model_infer.mode_backend.eplb.runtime_manager import EPLBManager self.eplb_manager = EPLBManager(self.model, self.args.eplb_rebalance_count) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py similarity index 95% rename from lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py rename to lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py index 6102158ebf..8168523e82 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/transfer_planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py @@ -7,8 +7,8 @@ from lightllm.utils.log_utils import init_logger -from .planner import ExpertPlacement -from .transfer import EPLBTransferInfo, build_transfer_plan +from .expert_transfer import EPLBTransferInfo, build_transfer_plan +from .placement_planner import ExpertPlacement logger = init_logger(__name__) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement.py b/lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py similarity index 100% rename from lightllm/server/router/model_infer/mode_backend/eplb/placement.py rename to lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py similarity index 100% rename from lightllm/server/router/model_infer/mode_backend/eplb/transfer.py rename to lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/plan.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py similarity index 96% rename from lightllm/server/router/model_infer/mode_backend/eplb/plan.py rename to lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py index 216c4cba59..8f3cf845bf 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/plan.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py @@ -9,7 +9,7 @@ from lightllm.utils.log_utils import init_logger -from .planner import EPLBPlanner, ExpertPlacement +from .placement_planner import EPLBPlanner, ExpertPlacement logger = init_logger(__name__) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement_planner.py similarity index 100% rename from lightllm/server/router/model_infer/mode_backend/eplb/planner.py rename to lightllm/server/router/model_infer/mode_backend/eplb/placement_planner.py diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py similarity index 98% rename from lightllm/server/router/model_infer/mode_backend/eplb/manager.py rename to lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 611eed1810..467197d654 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -18,14 +18,14 @@ from lightllm.utils.log_utils import init_logger from lightllm.utils.shm_port_args import get_shm_port_args -from .placement import build_logical_to_physical_map -from .plan import EPLBPlanTask -from .planner import EPLBPlanner, ExpertPlacement, GreedyEPLBPlanner -from .transfer import ( +from .async_transfer_planner import EPLBTransferPlanner +from .expert_placement import build_logical_to_physical_map +from .expert_transfer import ( EPLBTransferInfo, PinnedMemoryEPLBTransfer, ) -from .transfer_planner import EPLBTransferPlanner +from .placement_plan_task import EPLBPlanTask +from .placement_planner import EPLBPlanner, ExpertPlacement, GreedyEPLBPlanner logger = init_logger(__name__) EPLB_EXPERT_ALIGNMENT = 128 diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index bf60ddf7ba..c3307a0fb3 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -6,12 +6,12 @@ import pytest import torch -from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( +from lightllm.server.router.model_infer.mode_backend.eplb.expert_placement import ( build_initial_local_expert_ids, build_logical_to_physical_map, build_logical_to_physical_maps_for_layers, ) -from lightllm.server.router.model_infer.mode_backend.eplb.planner import ( +from lightllm.server.router.model_infer.mode_backend.eplb.placement_planner import ( EPLBPlanner, GreedyEPLBPlanner, ) @@ -19,16 +19,16 @@ from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.server.router.model_infer.infer_batch import g_infer_context from lightllm.server.router.model_infer.mode_backend.eplb import ( - manager as manager_module, + runtime_manager as manager_module, ) from lightllm.server.router.model_infer.mode_backend.eplb import ( - plan as plan_module, + placement_plan_task as plan_module, ) from lightllm.server.router.model_infer.mode_backend.eplb import ( - transfer as transfer_module, + expert_transfer as transfer_module, ) from lightllm.server.router.model_infer.mode_backend.eplb import ( - transfer_planner as transfer_planner_module, + async_transfer_planner as transfer_planner_module, ) from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( deepgemm_impl as deepgemm_module, @@ -45,7 +45,7 @@ fused_moe_weight as fused_weight_module, ) from lightllm.common.eplb_utils import extract_eplb_expert_tensors -from lightllm.server.router.model_infer.mode_backend.eplb.transfer import ( +from lightllm.server.router.model_infer.mode_backend.eplb.expert_transfer import ( EPLBTransferInfo, ExpertTensorBuffer, PinnedMemoryEPLBTransfer, diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 17e63b3eb6..da3387498d 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -11,7 +11,7 @@ import torch.distributed as dist import torch.multiprocessing as mp -from lightllm.server.router.model_infer.mode_backend.eplb.transfer import ( +from lightllm.server.router.model_infer.mode_backend.eplb.expert_transfer import ( PinnedMemoryEPLBTransfer, TransferStatus, build_transfer_plan, From b81ccbb54bd2d484a016a0fadbc2f21fc71e5d5e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 05:16:20 +0000 Subject: [PATCH 51/72] commit EPLB updates on overlap stream --- .../mode_backend/eplb/runtime_manager.py | 21 ++- unit_tests/common/fused_moe/test_eplb.py | 47 +++++- .../fused_moe/test_eplb_transfer_gpu.py | 140 ++++++++++++++++++ 3 files changed, 194 insertions(+), 14 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 467197d654..48e52056c6 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -370,11 +370,14 @@ def _poll_transfer_batch(self) -> None: # 始终一致;只有 destination rank 会额外写入实际专家权重。 from lightllm.server.router.model_infer.infer_batch import g_infer_context - torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - for transfer_info in self.active_transfer_batch: - self._commit_transfer(transfer_info) - for layer_index in {transfer_info.layer_index for transfer_info in self.active_transfer_batch}: - self._publish_layer_metadata(layer_index) + # 专家权重和路由 metadata 都由 overlap stream 上的 MoE forward + # 读取。将整批写操作排到同一条 stream,便可自然等待此前的 forward, + # 并保证后续 forward 只能看到完整提交后的权重与 metadata。 + with torch.cuda.stream(g_infer_context.get_overlap_stream()): + for transfer_info in self.active_transfer_batch: + self._commit_transfer(transfer_info) + for layer_index in {transfer_info.layer_index for transfer_info in self.active_transfer_batch}: + self._publish_layer_metadata(layer_index) del self.active_transfers del self.active_transfer_batch @@ -389,7 +392,10 @@ def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: ) assert active_transfer is not None, "EPLB destination rank has no matching completed transfer" for tensor_buffer in active_transfer.tensor_buffers: - tensor_buffer.live_tensor[transfer_info.dest_local_expert_index].copy_(tensor_buffer.pinned_row) + tensor_buffer.live_tensor[transfer_info.dest_local_expert_index].copy_( + tensor_buffer.pinned_row, + non_blocking=True, + ) layer_index = transfer_info.layer_index layer_impl = self._eplb_impls[layer_index] @@ -411,8 +417,9 @@ def _publish_layer_metadata(self, layer_index: int) -> None: current_rank=self.global_rank, ), dtype=torch.int32, + pin_memory=True, ) - layer_impl.logical_to_physical_map.copy_(logical_to_physical_map) + layer_impl.logical_to_physical_map.copy_(logical_to_physical_map, non_blocking=True) def _publish_expert_load_metric(self, local_load: torch.Tensor) -> None: if self.global_rank != 0: diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index c3307a0fb3..89708495d8 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -1390,7 +1390,17 @@ def __init__(self, offset, scale=True, zero_point=True): ] -def test_manager_commits_transfer_rows_and_metadata(): +def test_manager_commits_transfer_rows_and_metadata(monkeypatch): + original_copy = torch.Tensor.copy_ + non_blocking_values = [] + copy_sources = [] + + def record_copy(tensor, source, non_blocking=False): + non_blocking_values.append(non_blocking) + copy_sources.append(source) + return original_copy(tensor, source, non_blocking=non_blocking) + + monkeypatch.setattr(torch.Tensor, "copy_", record_copy) live = torch.arange(20).reshape(5, 4) original_primary = live[:3].clone() local_expert_ids = [0, 1, 2, 3, 2] @@ -1438,6 +1448,8 @@ def test_manager_commits_transfer_rows_and_metadata(): assert torch.equal(live[4], torch.full((4,), -5)) assert local_expert_ids == [0, 1, 2, 4, 5] assert torch.equal(logical_to_physical_map, expected_metadata) + assert non_blocking_values == [True, True, True] + assert copy_sources[-1].is_pinned() def test_manager_transfers_only_local_tasks_and_gathers_global_status(monkeypatch): @@ -1476,16 +1488,37 @@ def is_finished(self): manager.max_rebalance_count = -1 manager.completed_rebalance_count = 0 committed = [] + active_streams = [] cleared_route_counters = [] - manager._commit_transfer = committed.append - manager._publish_layer_metadata = lambda _layer_index: None + + def commit_transfer(transfer_info): + assert active_streams == [overlap_stream] + committed.append(transfer_info) + + def publish_layer_metadata(_layer_index): + assert active_streams == [overlap_stream] + + manager._commit_transfer = commit_transfer + manager._publish_layer_metadata = publish_layer_metadata manager._clear_route_counters = lambda: cleared_route_counters.append(True) - waits = [] + used_streams = [] overlap_stream = object() + + class StreamContext: + def __init__(self, stream): + self.stream = stream + + def __enter__(self): + used_streams.append(self.stream) + active_streams.append(self.stream) + + def __exit__(self, *_args): + active_streams.pop() + monkeypatch.setattr( manager_module.torch.cuda, - "current_stream", - lambda: SimpleNamespace(wait_stream=lambda stream: waits.append(stream)), + "stream", + StreamContext, ) monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) monkeypatch.setattr( @@ -1539,7 +1572,7 @@ def all_gather_object(output, local_state, **_kwargs): assert not hasattr(manager, "rebalance_started_at") assert cleared_route_counters == [True] assert local_states == [False, True, True] - assert waits == [overlap_stream, overlap_stream] + assert used_streams == [overlap_stream, overlap_stream] def test_manager_returns_to_collecting_after_reaching_rebalance_limit(): diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index da3387498d..99c21e721e 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -12,6 +12,7 @@ import torch.multiprocessing as mp from lightllm.server.router.model_infer.mode_backend.eplb.expert_transfer import ( + EPLBTransferInfo, PinnedMemoryEPLBTransfer, TransferStatus, build_transfer_plan, @@ -47,6 +48,29 @@ def _pack(logical_ids, layer_index, offset): return _Pack(values, scales) +class _StressFakeWeight: + """为并发传输压力测试生成可识别 source rank 的专家权重。""" + + def __init__(self, rank, layer_index, num_logical_experts): + self.layer_num_ = layer_index + logical_ids = list(range(num_logical_experts)) + self.fuse_moe_impl = SimpleNamespace(local_logics_expert_ids_list=logical_ids) + self.w13 = self._pack(rank, logical_ids, layer_index, 0) + self.w2 = self._pack(rank, logical_ids, layer_index, 100) + + @staticmethod + def _pack(rank, logical_ids, layer_index, offset): + # rank、layer、expert 和 tensor 类型都编码进数值;任何 recv 串包都会 + # 在目标 rank 的逐任务校验中表现为数值不一致。 + values = torch.tensor( + [[rank * 100_000 + layer_index * 1_000 + expert + offset] for expert in logical_ids], + dtype=torch.float32, + device="cuda", + ) + scales = values + 0.5 + return _Pack(values, scales) + + def _free_port(): sock = socket.socket() sock.bind(("127.0.0.1", 0)) @@ -66,6 +90,19 @@ def _wait_for_transfer(transfer, control_group): raise TimeoutError("EPLB transfer worker did not finish globally") +def _wait_for_all_transfers(transfers, control_group): + """等待每个 rank 参与的全部并发传输完成。""" + deadline = time.monotonic() + 120 + while time.monotonic() < deadline: + local_finished = all(transfer.status is TransferStatus.SUCCEEDED for transfer in transfers) + globally_finished = torch.tensor([int(local_finished)], dtype=torch.int32) + dist.all_reduce(globally_finished, op=dist.ReduceOp.MIN, group=control_group) + if int(globally_finished.item()) == 1: + return + time.sleep(0.001) + raise TimeoutError("concurrent EPLB transfer workers did not finish globally") + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_pytorch_keeps_unreferenced_pinned_source_alive_until_async_copy_finishes(): """验证 pinned allocator 不会提前复用仍被异步 H2D 读取的内存。 @@ -177,3 +214,106 @@ def _worker(rank, port): ) def test_eplb_pinned_memory_transfer_two_gpu_correctness(): mp.spawn(_worker, args=(_free_port(),), nprocs=2, join=True) + + +def _many_concurrent_p2p_worker(rank, port): + """同时运行大量、重复 rank 对的 PinnedMemoryEPLBTransfer。""" + world_size = 4 + num_layers = 8 + num_logical_experts = 32 + transfers_per_rank_pair_per_layer = 8 + + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = str(port) + torch.cuda.set_device(rank) + dist.init_process_group("gloo", rank=rank, world_size=world_size) + control_group = dist.new_group(list(range(world_size)), backend="gloo") + transfer_group = dist.new_group(list(range(world_size)), backend="gloo") + + try: + weights = [_StressFakeWeight(rank, layer_index, num_logical_experts) for layer_index in range(num_layers)] + + # 每层覆盖全部 12 个有向 rank 对,每个 rank 对重复 8 次。任务 identity + # 中的 layer、expert 和目标槽位不同,因此应该获得独立的 Gloo tag; + # source/destination rank 对则会被大量重复使用。 + transfer_infos = [] + for layer_index in range(num_layers): + for source_rank in range(world_size): + for dest_rank in range(world_size): + if source_rank == dest_rank: + continue + source_peers = [peer_rank for peer_rank in range(world_size) if peer_rank != dest_rank] + source_peer_index = source_peers.index(source_rank) + for repeat_index in range(transfers_per_rank_pair_per_layer): + expert_id = source_rank * transfers_per_rank_pair_per_layer + repeat_index + dest_local_expert_index = source_peer_index * transfers_per_rank_pair_per_layer + repeat_index + transfer_infos.append( + EPLBTransferInfo( + source_rank=source_rank, + layer_index=layer_index, + source_logical_expert_id=expert_id, + dest_rank=dest_rank, + dest_local_expert_index=dest_local_expert_index, + ) + ) + + local_transfer_infos = [ + transfer_info + for transfer_info in transfer_infos + if rank in (transfer_info.source_rank, transfer_info.dest_rank) + ] + transfers = [ + PinnedMemoryEPLBTransfer(weights, transfer_group, rank, transfer_info) + for transfer_info in local_transfer_infos + ] + assert len(transfer_infos) == 768 + assert len(transfers) == 384 + + # 同一个 source/destination 对上的并发消息必须具有不同 tag,否则不同 + # 专家或张量可能被错误匹配。不同 rank 对可以安全复用相同整数 tag。 + message_keys = [] + for transfer in transfers: + transfer_info = transfer.transfer_info + for tensor_buffer in transfer.tensor_buffers: + message_keys.append( + ( + transfer_info.source_rank, + transfer_info.dest_rank, + transfer._build_p2p_message_tag(tensor_buffer.name), + ) + ) + assert len(message_keys) == len(set(message_keys)) + + dist.barrier(group=control_group) + for transfer in transfers: + transfer.start() + _wait_for_all_transfers(transfers, control_group) + + destination_transfers = [transfer for transfer in transfers if transfer.transfer_info.dest_rank == rank] + assert len(destination_transfers) == 192 + for transfer in destination_transfers: + transfer_info = transfer.transfer_info + expected_w13 = ( + transfer_info.source_rank * 100_000 + + transfer_info.layer_index * 1_000 + + transfer_info.source_logical_expert_id + ) + expected_values = [expected_w13, expected_w13 + 0.5, expected_w13 + 100, expected_w13 + 100.5] + assert [buffer.name for buffer in transfer.tensor_buffers] == [ + "w13.weight", + "w13.weight_scale", + "w2.weight", + "w2.weight_scale", + ] + for tensor_buffer, expected_value in zip(transfer.tensor_buffers, expected_values): + assert torch.all(tensor_buffer.pinned_row == expected_value) + finally: + dist.destroy_process_group() + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < 4, + reason="requires four CUDA GPUs", +) +def test_eplb_pinned_memory_transfer_four_gpu_many_concurrent_p2p(): + mp.spawn(_many_concurrent_p2p_worker, args=(_free_port(),), nprocs=4, join=True) From 89b4474a3baa2861c2e327f064f24074e433097d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 05:21:56 +0000 Subject: [PATCH 52/72] remove unused EPLB layer map helper --- .../mode_backend/eplb/expert_placement.py | 49 ++----------------- unit_tests/common/fused_moe/test_eplb.py | 26 ---------- 2 files changed, 5 insertions(+), 70 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py b/lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py index f41b757f29..2eec912ff9 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py @@ -1,9 +1,9 @@ -"""构建完整专家布局及其紧凑路由元数据。 +"""构建单层专家布局及其紧凑路由元数据。 -本模块统一使用 -``[layer][rank][local physical expert] -> logical expert`` 表示专家布局; -处理单层布局的函数会省略 layer 维。初始布局按主专家和冗余专家构建; -EPLB 开始运行后,全部物理槽位都可以重新分配。 +本模块统一使用 ``[rank][local physical expert] -> logical expert`` 表示一层 +完整的专家布局。初始布局按主专家和冗余专家构建;EPLB 开始运行后,全部 +物理槽位都可以重新分配。多层调用方应逐层调用这里的单层函数,避免维护 +语义重复的批量封装。 """ @@ -177,42 +177,3 @@ def _build_routing_row( # 阶段 3:第 0 列保存 kernel 参与 hash 的有效副本数;第 1 列标记是否 # 存在本地副本;后续列保存按本地优先顺序排列的 physical IDs 和 -1 padding。 return [num_valid_replicas, int(has_local_replica), *routing_slots] - - -def build_logical_to_physical_maps_for_layers( - rank_to_logic_expert_ids_by_layer: list[list[list[int]]], - num_logical_experts: int, - current_rank: int, -) -> list[list[list[int]]]: - """逐层构建完整布局对应的 logical-to-physical 路由表。 - - ``rank_to_logic_expert_ids_by_layer`` 使用 EPLB 的完整布局格式: - ``[layer][rank][local physical expert] -> logical expert``。其中每个 layer - 必须包含相同数量的 rank,同一层内每个 rank 必须具有相同数量的物理 - 专家槽位;布局同时包含初始主专家和全部冗余专家。 - - 返回值的 shape 为 ``[layer][logical expert][metadata]``。每个 logical - expert 的 metadata 行格式与 :func:`build_logical_to_physical_map` 一致: - - * 第 0 项是该逻辑专家当前有效的物理副本数; - * 第 1 项标记 ``current_rank`` 是否持有本地副本; - * 第 2 项起保存可参与路由的全局 physical expert ID; - * 如果存在本地副本,本地 physical ID 固定排在第一个路由槽; - * 有效副本之后的固定宽度空槽使用 ``-1`` 填充。 - - 各层相互独立,并按输入中的 layer 顺序逐层调用单层构建函数。这样初始化、 - 在线重排和批量 metadata 发布共享完全相同的副本排序及打包规则,不会出现 - 单层接口和多层接口语义偏差。 - - 本函数只构建普通 CPU list,不创建或搬运 Tensor。调用方需要更新 GPU - 路由 metadata 时,应在通信或提交边界统一转换为对应 dtype/device 的 - ``torch.Tensor``。 - """ - return [ - build_logical_to_physical_map( - layer_placement, - num_logical_experts, - current_rank=current_rank, - ) - for layer_placement in rank_to_logic_expert_ids_by_layer - ] diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 89708495d8..0252249b8b 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -9,7 +9,6 @@ from lightllm.server.router.model_infer.mode_backend.eplb.expert_placement import ( build_initial_local_expert_ids, build_logical_to_physical_map, - build_logical_to_physical_maps_for_layers, ) from lightllm.server.router.model_infer.mode_backend.eplb.placement_planner import ( EPLBPlanner, @@ -646,31 +645,6 @@ def test_current_rank_stably_moves_all_local_physical_ids_to_front(): assert rank1_map[0] == [4, 1, 3, 0, 1, 5] -@pytest.mark.parametrize("current_rank", [0, 1, 2, 3]) -def test_logical_to_physical_maps_for_layers_match_single_layer_api(current_rank): - placements_by_layer = [ - _rank_to_logic_expert_ids([[4, 5], [0, 1], [0, 1], [2, 3]], 8), - _rank_to_logic_expert_ids([[6, 7], [0, 1], [0, 1], [2, 3]], 8), - _rank_to_logic_expert_ids([[4, 5], [0, 1], [0, 1], [2, 3]], 8), - ] - - maps_by_layer = build_logical_to_physical_maps_for_layers( - placements_by_layer, - num_logical_experts=8, - current_rank=current_rank, - ) - expected_maps = [ - build_logical_to_physical_map( - layer_placement, - num_logical_experts=8, - current_rank=current_rank, - ) - for layer_placement in placements_by_layer - ] - - assert maps_by_layer == expected_maps - - def test_transfer_plan_respects_explicit_target_slots(): current = [[0, 1, 4, 5], [2, 3, 6, 7], [4, 5, 0, 1], [6, 7, 2, 3]] target = [[0, 1, 5, 4], [2, 3, 7, 6], [4, 5, 1, 0], [6, 7, 3, 2]] From d7a510cc1eed5def0b2e6e720466bd101e9743f5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 05:53:26 +0000 Subject: [PATCH 53/72] refactor(eplb): organize placement modules --- .../fused_moe/impl/deepgemm_impl.py | 2 +- .../eplb/async_transfer_planner.py | 2 +- .../mode_backend/eplb/placement/__init__.py | 18 +++++++ .../greedy.py} | 43 +++++---------- .../mode_backend/eplb/placement/initial.py | 39 ++++++++++++++ .../mode_backend/eplb/placement/planner.py | 17 ++++++ .../routing.py} | 52 ++++--------------- .../mode_backend/eplb/placement/types.py | 15 ++++++ .../mode_backend/eplb/placement_plan_task.py | 2 +- .../mode_backend/eplb/runtime_manager.py | 8 ++- unit_tests/common/fused_moe/test_eplb.py | 8 ++- 11 files changed, 124 insertions(+), 82 deletions(-) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py rename lightllm/server/router/model_infer/mode_backend/eplb/{placement_planner.py => placement/greedy.py} (94%) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/placement/initial.py create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py rename lightllm/server/router/model_infer/mode_backend/eplb/{expert_placement.py => placement/routing.py} (75%) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/placement/types.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 7da12de0a0..0c9a6e3b3a 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 @@ -1,7 +1,7 @@ import torch from typing import Optional, Tuple, Any from .base_impl import FuseMoeBaseImpl -from lightllm.server.router.model_infer.mode_backend.eplb.expert_placement import ( +from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( build_initial_local_expert_ids, build_logical_to_physical_map, ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py index 8168523e82..3e3746a63c 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py @@ -8,7 +8,7 @@ from lightllm.utils.log_utils import init_logger from .expert_transfer import EPLBTransferInfo, build_transfer_plan -from .placement_planner import ExpertPlacement +from .placement import ExpertPlacement logger = init_logger(__name__) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py new file mode 100644 index 0000000000..d63c414f14 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py @@ -0,0 +1,18 @@ +"""Expert placement construction, routing metadata, and planning APIs.""" + +from .types import ExpertPlacement, LayerPlacement, LogicalExpertLoad, LogicalToPhysicalMap +from .planner import EPLBPlanner +from .initial import build_initial_local_expert_ids +from .routing import build_logical_to_physical_map +from .greedy import GreedyEPLBPlanner + +__all__ = [ + "EPLBPlanner", + "ExpertPlacement", + "GreedyEPLBPlanner", + "LayerPlacement", + "LogicalExpertLoad", + "LogicalToPhysicalMap", + "build_initial_local_expert_ids", + "build_logical_to_physical_map", +] diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py similarity index 94% rename from lightllm/server/router/model_infer/mode_backend/eplb/placement_planner.py rename to lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py index 44cbd86c9d..3a68ba2dc3 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement_planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py @@ -1,33 +1,15 @@ -"""使用纯 Python 实现 EPLB 专家布局规划。 +"""使用纯 Python 实现贪心 EPLB 专家布局规划。 规划器有意使用嵌套 list,而不是 Tensor。Tensor 转换仅发生在 manager 的 分布式通信和迁移边界;规划模块不依赖 Tensor,更易于阅读、测试和替换算法。 """ import heapq -from abc import ABC, abstractmethod from math import ceil -from typing import List, Tuple +from typing import List - -# [layer][logical expert] -LogicalExpertLoad = List[List[float]] -# [layer][rank][local physical expert] -> logical expert -ExpertPlacement = List[List[List[int]]] -# (logical expert, replica count, aligned load per replica) -ExpertReplicaGroup = Tuple[int, int, float] - - -class EPLBPlanner(ABC): - """专家布局规划接口。""" - - @abstractmethod - def plan( - self, - logical_expert_load: LogicalExpertLoad, - current_placement: ExpertPlacement, - ) -> ExpertPlacement: - """返回完整的 ``[layer][rank][local physical expert]`` 专家布局。""" +from .planner import EPLBPlanner +from .types import ExpertPlacement, ExpertReplicaGroup, LayerPlacement, LogicalExpertLoad class GreedyEPLBPlanner(EPLBPlanner): @@ -180,8 +162,8 @@ def plan( def _plan_layer( self, logical_load: List[float], - current_placement: List[List[int]], - ) -> List[List[int]]: + current_placement: LayerPlacement, + ) -> LayerPlacement: """完成单层副本分配、rank 排布和物理槽位复用。""" # 阶段 1:选出最热的 R 个专家,并为每个 rank 固定预留它们的副本。 redundant_experts = self._select_redundant_experts(logical_load) @@ -219,7 +201,10 @@ def _build_remaining_expert_groups( for _ in range(self.num_redundant_experts_per_rank): expert = min( (expert for expert in remaining_experts if replica_counts[expert] < self.world_size), - key=lambda expert: (-logical_load[expert] / replica_counts[expert], expert), + key=lambda expert: ( + -logical_load[expert] / replica_counts[expert], + expert, + ), ) replica_counts[expert] += 1 @@ -241,7 +226,7 @@ def _distribute_remaining_experts( self, redundant_experts: List[int], expert_groups: List[ExpertReplicaGroup], - ) -> List[List[int]]: + ) -> LayerPlacement: """先平铺多副本专家,再按当前 rank 负载分配单副本专家。""" placement = [list(redundant_experts) for _ in range(self.world_size)] @@ -293,9 +278,9 @@ def _distribute_remaining_experts( def _reuse_current_slots( self, - candidate_placement: List[List[int]], - current_placement: List[List[int]], - ) -> List[List[int]]: + candidate_placement: LayerPlacement, + current_placement: LayerPlacement, + ) -> LayerPlacement: """贪心匹配候选 rank,并让共同专家尽量复用当前物理槽位。""" # 步骤 1:准备所有尚未匹配的候选行。 # diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/initial.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/initial.py new file mode 100644 index 0000000000..626b7caa15 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/initial.py @@ -0,0 +1,39 @@ +"""Construct the deterministic expert placement used during model loading.""" + +from .types import LayerPlacement + + +def build_initial_local_expert_ids( + num_logical_experts: int, + num_ranks: int, + num_redundant_experts_per_rank: int, +) -> LayerPlacement: + """构建每个 rank 初始持有的完整 logical expert ID 列表。 + + 每个 rank 先持有连续划分得到的主专家,再按 rank 顺序选择不属于 + 本 rank 的专家作为默认冗余副本。这里仅负责生成 Python 列表;调用方如果要参与 tensor + 运算,需要自行转换为 ``torch.Tensor``。 + + 例如 ``num_logical_experts=8``、``num_ranks=4``、每个 rank 有 2 个 + 额外槽时,每个 rank 分到 2 个主专家,结果为: + + ``[[0, 1, 2, 3], [2, 3, 4, 5], [4, 5, 6, 7], [6, 7, 0, 1]]`` + + 其中每行前两个值是主专家,后两个值是已在初始加载阶段就可用的冗余副本。 + """ + assert num_logical_experts % num_ranks == 0 + num_experts_per_rank = num_logical_experts // num_ranks + assert 0 <= num_redundant_experts_per_rank <= num_logical_experts - num_experts_per_rank + + local_expert_ids_by_rank = [] + for rank in range(num_ranks): + first_expert_id = rank * num_experts_per_rank + local_expert_ids = list(range(first_expert_id, first_expert_id + num_experts_per_rank)) + first_redundant_expert_id = ((rank + 1) * num_experts_per_rank) % num_logical_experts + local_expert_ids.extend( + (first_redundant_expert_id + offset) % num_logical_experts + for offset in range(num_redundant_experts_per_rank) + ) + local_expert_ids_by_rank.append(local_expert_ids) + + return local_expert_ids_by_rank diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py new file mode 100644 index 0000000000..866fe8f4aa --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py @@ -0,0 +1,17 @@ +"""Abstract interface implemented by EPLB placement planners.""" + +from abc import ABC, abstractmethod + +from .types import ExpertPlacement, LogicalExpertLoad + + +class EPLBPlanner(ABC): + """根据逻辑专家负载生成完整物理布局。""" + + @abstractmethod + def plan( + self, + logical_expert_load: LogicalExpertLoad, + current_placement: ExpertPlacement, + ) -> ExpertPlacement: + """返回 ``[layer][rank][local physical expert]`` 专家布局。""" diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/routing.py similarity index 75% rename from lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py rename to lightllm/server/router/model_infer/mode_backend/eplb/placement/routing.py index 2eec912ff9..d03ce297cb 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/expert_placement.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/routing.py @@ -1,53 +1,19 @@ -"""构建单层专家布局及其紧凑路由元数据。 +"""Build compact logical-to-physical routing metadata for one MoE layer. -本模块统一使用 ``[rank][local physical expert] -> logical expert`` 表示一层 -完整的专家布局。初始布局按主专家和冗余专家构建;EPLB 开始运行后,全部 -物理槽位都可以重新分配。多层调用方应逐层调用这里的单层函数,避免维护 -语义重复的批量封装。 +The input layout uses ``[rank][local physical expert] -> logical expert``. +Initial placement construction lives in :mod:`.initial`; this module only +inverts an existing placement into the fixed-width rows consumed by the +EPLB routing kernel. """ - -def build_initial_local_expert_ids( - num_logical_experts: int, - num_ranks: int, - num_redundant_experts_per_rank: int, -) -> list[list[int]]: - """构建每个 rank 初始持有的完整 logical expert ID 列表。 - - 每个 rank 先持有连续划分得到的主专家,再按 rank 顺序选择不属于 - 本 rank 的专家作为默认冗余副本。这里仅负责生成 Python 列表;调用方如果要参与 tensor - 运算,需要自行转换为 ``torch.Tensor``。 - - 例如 ``num_logical_experts=8``、``num_ranks=4``、每个 rank 有 2 个 - 额外槽时,每个 rank 分到 2 个主专家,结果为: - - ``[[0, 1, 2, 3], [2, 3, 4, 5], [4, 5, 6, 7], [6, 7, 0, 1]]`` - - 其中每行前两个值是主专家,后两个值是已在初始加载阶段就可用的冗余副本。 - """ - assert num_logical_experts % num_ranks == 0 - num_experts_per_rank = num_logical_experts // num_ranks - assert 0 <= num_redundant_experts_per_rank <= num_logical_experts - num_experts_per_rank - - local_expert_ids_by_rank = [] - for rank in range(num_ranks): - first_expert_id = rank * num_experts_per_rank - local_expert_ids = list(range(first_expert_id, first_expert_id + num_experts_per_rank)) - first_redundant_expert_id = ((rank + 1) * num_experts_per_rank) % num_logical_experts - local_expert_ids.extend( - (first_redundant_expert_id + offset) % num_logical_experts - for offset in range(num_redundant_experts_per_rank) - ) - local_expert_ids_by_rank.append(local_expert_ids) - - return local_expert_ids_by_rank +from .types import LayerPlacement, LogicalToPhysicalMap def build_logical_to_physical_map( - rank_to_logic_expert_ids: list[list[int]], + rank_to_logic_expert_ids: LayerPlacement, num_logical_experts: int, current_rank: int, -) -> list[list[int]]: +) -> LogicalToPhysicalMap: """使用普通 CPU list 构建单层 logical 到 physical expert 的路由表。 ``rank_to_logic_expert_ids`` 的 shape 为 @@ -110,7 +76,7 @@ def build_logical_to_physical_map( def _collect_physical_ids_by_logical_expert( - rank_to_logic_expert_ids: list[list[int]], + rank_to_logic_expert_ids: LayerPlacement, num_logical_experts: int, ) -> list[list[int]]: """将完整物理布局反转为每个 logical expert 对应的物理槽位。 diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/types.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/types.py new file mode 100644 index 0000000000..b482f4bcdd --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/types.py @@ -0,0 +1,15 @@ +"""Shared type aliases for EPLB placement planning.""" + +from typing import List, Tuple + + +# [layer][logical expert] +LogicalExpertLoad = List[List[float]] +# [rank][local physical expert] -> logical expert +LayerPlacement = List[List[int]] +# [layer][rank][local physical expert] -> logical expert +ExpertPlacement = List[LayerPlacement] +# [logical expert][replica metadata] +LogicalToPhysicalMap = List[List[int]] +# (logical expert, replica count, aligned load per replica) +ExpertReplicaGroup = Tuple[int, int, float] diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py index 8f3cf845bf..8fc80d8574 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py @@ -9,7 +9,7 @@ from lightllm.utils.log_utils import init_logger -from .placement_planner import EPLBPlanner, ExpertPlacement +from .placement import EPLBPlanner, ExpertPlacement logger = init_logger(__name__) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 48e52056c6..4fcdd20f67 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -19,13 +19,17 @@ from lightllm.utils.shm_port_args import get_shm_port_args from .async_transfer_planner import EPLBTransferPlanner -from .expert_placement import build_logical_to_physical_map from .expert_transfer import ( EPLBTransferInfo, PinnedMemoryEPLBTransfer, ) +from .placement import ( + EPLBPlanner, + ExpertPlacement, + GreedyEPLBPlanner, + build_logical_to_physical_map, +) from .placement_plan_task import EPLBPlanTask -from .placement_planner import EPLBPlanner, ExpertPlacement, GreedyEPLBPlanner logger = init_logger(__name__) EPLB_EXPERT_ALIGNMENT = 128 diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 0252249b8b..2fab575816 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -6,13 +6,11 @@ import pytest import torch -from lightllm.server.router.model_infer.mode_backend.eplb.expert_placement import ( - build_initial_local_expert_ids, - build_logical_to_physical_map, -) -from lightllm.server.router.model_infer.mode_backend.eplb.placement_planner import ( +from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( EPLBPlanner, GreedyEPLBPlanner, + build_initial_local_expert_ids, + build_logical_to_physical_map, ) from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs From ae80f60dffa24944dfbac497cfd033894b447929 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 05:57:17 +0000 Subject: [PATCH 54/72] fix(eplb): reject redundant experts with RL --- lightllm/server/api_start.py | 2 ++ unit_tests/server/test_api_start_eplb.py | 24 ++++++++++++++++++++++++ 2 files changed, 26 insertions(+) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 94eb52bfa4..29ac919f37 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -166,6 +166,8 @@ def _launch_subprocesses(args: StartArgs): if args.eplb_num_redundant_experts_per_rank > 0: assert args.enable_ep_moe, "EPLB requires --enable_ep_moe" assert not args.enable_prefill_cudagraph, "EPLB does not support --enable_prefill_cudagraph" + # TODO: Support EPLB redundant experts together with RL after their runtime state updates are coordinated. + assert not args.enable_rl, "EPLB redundant experts do not support --enable_rl" # EPLB updates expert weights in place, but SM100 Mega-MoE caches transformed weights by tensor data_ptr. assert not is_sm100_gpu(), "EPLB does not support SM100" diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py index a1da20b1ba..2604fbef1a 100644 --- a/unit_tests/server/test_api_start_eplb.py +++ b/unit_tests/server/test_api_start_eplb.py @@ -86,6 +86,30 @@ def test_eplb_prefill_cudagraph_is_rejected_before_starting_subprocesses(monkeyp api_start._launch_subprocesses(args) +def test_eplb_redundant_experts_cannot_be_combined_with_rl(monkeypatch): + args = StartArgs( + enable_ep_moe=True, + enable_rl=True, + eplb_num_redundant_experts_per_rank=2, + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) + monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) + monkeypatch.setattr( + api_start.process_manager, + "start_submodule_processes", + lambda *args, **kwargs: pytest.fail("subprocess startup must not be reached"), + ) + + with pytest.raises(AssertionError, match="EPLB redundant experts do not support --enable_rl"): + api_start._launch_subprocesses(args) + + def test_eplb_mtp_combination_is_not_rejected_before_starting_subprocesses(monkeypatch): args = StartArgs( model_dir="test-model", From 8ddb29bf3bb8b5945fe891e67f12fc28c6c92a9b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 06:03:25 +0000 Subject: [PATCH 55/72] refactor(eplb): move weight utilities into runtime package --- .../router/model_infer/mode_backend/eplb}/eplb_utils.py | 2 +- .../router/model_infer/mode_backend/eplb/expert_transfer.py | 3 ++- unit_tests/common/fused_moe/test_eplb.py | 2 +- 3 files changed, 4 insertions(+), 3 deletions(-) rename lightllm/{common => server/router/model_infer/mode_backend/eplb}/eplb_utils.py (95%) diff --git a/lightllm/common/eplb_utils.py b/lightllm/server/router/model_infer/mode_backend/eplb/eplb_utils.py similarity index 95% rename from lightllm/common/eplb_utils.py rename to lightllm/server/router/model_infer/mode_backend/eplb/eplb_utils.py index 073847292a..b84afcb20d 100644 --- a/lightllm/common/eplb_utils.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/eplb_utils.py @@ -1,4 +1,4 @@ -"""传输与模型分析共用的轻量 EPLB 工具。""" +"""EPLB 专家权重提取工具。""" from typing import List, Optional, Protocol, Tuple diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py index de8f4ce04f..cc452cb5ad 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py @@ -11,9 +11,10 @@ import torch.distributed as dist from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight -from lightllm.common.eplb_utils import NamedTensor, extract_eplb_expert_tensors from lightllm.utils.log_utils import init_logger +from .eplb_utils import NamedTensor, extract_eplb_expert_tensors + logger = init_logger(__name__) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 2fab575816..4d3cbb3f2a 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -41,7 +41,7 @@ from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe import ( fused_moe_weight as fused_weight_module, ) -from lightllm.common.eplb_utils import extract_eplb_expert_tensors +from lightllm.server.router.model_infer.mode_backend.eplb.eplb_utils import extract_eplb_expert_tensors from lightllm.server.router.model_infer.mode_backend.eplb.expert_transfer import ( EPLBTransferInfo, ExpertTensorBuffer, From e3316cc22dc3867f40c5ac8c620cdf395582193c Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 06:16:43 +0000 Subject: [PATCH 56/72] fix(eplb): validate SM100 support in runtime manager --- lightllm/server/api_start.py | 3 --- .../model_infer/mode_backend/eplb/runtime_manager.py | 7 +++++++ unit_tests/common/fused_moe/test_eplb.py | 9 +++++++++ unit_tests/server/test_api_start_eplb.py | 1 - 4 files changed, 16 insertions(+), 4 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 29ac919f37..7c72a7cbbc 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -25,7 +25,6 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args -from lightllm.utils.device_utils import is_sm100_gpu logger = init_logger(__name__) @@ -168,8 +167,6 @@ def _launch_subprocesses(args: StartArgs): assert not args.enable_prefill_cudagraph, "EPLB does not support --enable_prefill_cudagraph" # TODO: Support EPLB redundant experts together with RL after their runtime state updates are coordinated. assert not args.enable_rl, "EPLB redundant experts do not support --enable_rl" - # EPLB updates expert weights in place, but SM100 Mega-MoE caches transformed weights by tensor data_ptr. - assert not is_sm100_gpu(), "EPLB does not support SM100" if args.enable_ep_moe: allowed_ep_prefill_att_backends = {"auto", "fa3", "triton", "flashqla"} diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 4fcdd20f67..2b24794aa4 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -14,6 +14,7 @@ get_global_rank, get_global_world_size, ) +from lightllm.utils.device_utils import is_sm100_gpu from lightllm.utils.envs_utils import get_eplb_step_interval from lightllm.utils.log_utils import init_logger from lightllm.utils.shm_port_args import get_shm_port_args @@ -67,6 +68,12 @@ class EPLBManager: """ def __init__(self, model: TpPartBaseModel, max_rebalance_count: int = 1) -> None: + # SM100 FP4 Mega-MoE 会将在线专家权重转换为独立的 kernel 布局,并使用源 tensor 的 data_ptr + # 作为 key 缓存这些转换后的副本。EPLB 通过原地 copy_ 替换专家行,只改变权重内容而不会改变 + # data_ptr,因此重平衡后 Mega-MoE 仍会读取旧的转换权重。在 EPLB 能够失效或更新该缓存前, + # 暂不支持 SM100。 + assert not is_sm100_gpu(), "EPLB does not support SM100" + weights: List[FusedMoeWeight] = _find_fused_moe_weights(model) assert weights, "EPLB requires at least one EP MoE layer" assert max_rebalance_count >= -1 diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 4d3cbb3f2a..c4583d9ea8 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -2136,6 +2136,7 @@ def synchronize(self): def test_manager_requires_more_than_one_rank(monkeypatch): + monkeypatch.setattr(manager_module, "is_sm100_gpu", lambda: False) monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [object()]) monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 1) @@ -2144,6 +2145,13 @@ def test_manager_requires_more_than_one_rank(monkeypatch): manager_module.EPLBManager(type("Model", (), {})()) +def test_manager_rejects_sm100_before_initialization(monkeypatch): + monkeypatch.setattr(manager_module, "is_sm100_gpu", lambda: True) + + with pytest.raises(AssertionError, match="EPLB does not support SM100"): + manager_module.EPLBManager(type("Model", (), {})()) + + def test_manager_clears_all_route_counters_on_overlap_stream(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) counters = [torch.tensor([1, 2]), torch.tensor([3, 4])] @@ -2182,6 +2190,7 @@ def test_manager_initializes_without_transfer_task(monkeypatch): )() groups = [object(), object()] new_group_calls = [] + monkeypatch.setattr(manager_module, "is_sm100_gpu", lambda: False) monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py index 2604fbef1a..74984b5b62 100644 --- a/unit_tests/server/test_api_start_eplb.py +++ b/unit_tests/server/test_api_start_eplb.py @@ -128,7 +128,6 @@ def test_eplb_mtp_combination_is_not_rejected_before_starting_subprocesses(monke monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) - monkeypatch.setattr(api_start, "is_sm100_gpu", lambda: False) monkeypatch.setattr(api_start, "auto_set_response_parsers", lambda args: None) monkeypatch.setattr(api_start, "auto_configure_allreduce_flags_from_args", lambda args: None) monkeypatch.setattr(api_start, "validate_ports", lambda ports: None) From 6da8ccac7106dadf09d1831a374ffc3ec681230b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 07:20:33 +0000 Subject: [PATCH 57/72] feat: persist EPLB placement configuration --- docs/CN/source/tutorial/api_server_args.rst | 13 +- docs/EN/source/tutorial/api_server_args.rst | 14 +- .../fused_moe/fused_moe_weight.py | 1 + .../meta_weights/fused_moe/impl/__init__.py | 4 + .../meta_weights/fused_moe/impl/base_impl.py | 2 + .../fused_moe/impl/deepgemm_impl.py | 24 +- lightllm/server/api_cli.py | 7 + lightllm/server/core/objs/start_args_type.py | 1 + .../model_infer/mode_backend/base_backend.py | 6 +- .../mode_backend/eplb/placement/__init__.py | 3 + .../mode_backend/eplb/placement/config.py | 230 ++++++++++++++++++ .../mode_backend/eplb/runtime_manager.py | 25 +- unit_tests/common/fused_moe/test_eplb.py | 69 +++++- .../fused_moe/test_eplb_placement_config.py | 202 +++++++++++++++ 14 files changed, 592 insertions(+), 9 deletions(-) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/placement/config.py create mode 100644 unit_tests/common/fused_moe/test_eplb_placement_config.py diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 6410394459..ad21c4a53d 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -733,13 +733,24 @@ PD 分离模式参数 * ``0``:不进行动态重排,仅使用初始化时的冗余布局; * 正整数:完成指定次数的重排后停止规划。 +.. option:: --eplb_config_path + + EPLB 布局 JSON 文件路径,默认值为 ``None``。指定后,LightLLM 会在初始化专家权重之前校验并 + 读取各层保存的布局,使服务启动后立即使用上一次优化得到的专家分配。同一个文件也作为输出: + 每次成功完成动态重排后,rank 0 会写回当前最新布局。 + + 如果文件不存在、JSON 无法解析、缺少模型层,或保存的专家拓扑与当前部署不匹配,LightLLM 会 + 记录 warning,并对受影响的层使用默认初始化布局。只有后续动态重排成功完成时,rank 0 才会把 + 新布局写入该路径。 + 以下示例为每个 EP rank 配置两个冗余专家:: python -m lightllm.server.api_server \ --model_dir /path/to/model \ --enable_ep_moe \ --eplb_num_redundant_experts_per_rank 2 \ - --eplb_rebalance_count 1 + --eplb_rebalance_count 1 \ + --eplb_config_path /path/to/eplb-placement.json MTP 多预测参数 -------------- diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 11c7a6a2ee..83a671f175 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -752,13 +752,25 @@ Expert Parallelism and EPLB Parameters * ``0`` disables dynamic rebalancing, leaving only the initial redundant placement active. * A positive value stops planning after that many completed rebalances. +.. option:: --eplb_config_path + + Path to an EPLB placement JSON file. The default is ``None``. When specified, LightLLM validates and loads + the saved per-layer placement before expert weights are initialized, so the service starts directly with the + previous optimized layout. The same file is updated with the latest layout after every successfully completed + rebalance. + + If the file does not exist, cannot be decoded, is missing a model layer, or does not match the current expert + topology, LightLLM logs a warning and uses the default initial placement for the affected layer. Rank 0 writes + a new layout to this path only after a dynamic rebalance completes successfully. + Example: enable EPLB with two redundant experts per EP rank:: python -m lightllm.server.api_server \ --model_dir /path/to/model \ --enable_ep_moe \ --eplb_num_redundant_experts_per_rank 2 \ - --eplb_rebalance_count 1 + --eplb_rebalance_count 1 \ + --eplb_config_path /path/to/eplb-placement.json MTP Multi-Prediction Parameters ------------------------------- 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 175b4313c7..c04a4022dc 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 @@ -62,6 +62,7 @@ def __init__( routed_scaling_factor=self.routed_scaling_factor, quant_method=self.quant_method, enable_ep_moe=self.enable_ep_moe, + layer_index=self.layer_num_, ) self._init_weight_partition() self.lock = threading.Lock() diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 16c6f17576..80a320cefa 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -1,3 +1,5 @@ +from typing import Optional + from lightllm.common.quantization.quantize_method import QuantizationMethod from .triton_impl import FuseMoeTriton from .marlin_impl import FuseMoeMarlin @@ -11,6 +13,7 @@ def create_fuse_moe_impl( routed_scaling_factor: float, quant_method: QuantizationMethod, enable_ep_moe: bool = False, + layer_index: Optional[int] = None, ): """创建持有自身路由运行态的 MoE 执行实现。 @@ -29,4 +32,5 @@ def create_fuse_moe_impl( num_fused_shared_experts=num_fused_shared_experts, routed_scaling_factor=routed_scaling_factor, quant_method=quant_method, + layer_index=layer_index, ) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py index fb4b1406f5..ca42db39cf 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py @@ -40,11 +40,13 @@ def __init__( num_fused_shared_experts: int, routed_scaling_factor: float, quant_method: QuantizationMethod, + layer_index: Optional[int] = None, ): self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self.routed_scaling_factor = routed_scaling_factor self.quant_method = quant_method + self.layer_index = layer_index def __call__( self, 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 0c9a6e3b3a..add55795ac 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 @@ -4,6 +4,7 @@ from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( build_initial_local_expert_ids, build_logical_to_physical_map, + load_layer_placement, ) from lightllm.distributed import dist_group_manager from lightllm.common.quantization.quantize_method import WeightPack @@ -45,15 +46,36 @@ def _init_eplb_runtime(self): world_size = get_global_world_size() assert self.n_routed_experts % world_size == 0 global_rank = get_global_rank() - self.num_redundant_experts_per_rank = get_env_start_args().eplb_num_redundant_experts_per_rank + start_args = get_env_start_args() + self.num_redundant_experts_per_rank = start_args.eplb_num_redundant_experts_per_rank if self.num_redundant_experts_per_rank > 0: self.num_total_physical_experts = self.n_routed_experts + world_size * self.num_redundant_experts_per_rank + + # 阶段 1:先构造确定性的默认布局。未指定配置文件,或配置读取、校验失败时, + # 后续权重初始化会继续使用这份布局。 initial_local_expert_ids_by_rank = build_initial_local_expert_ids( self.n_routed_experts, world_size, self.num_redundant_experts_per_rank, ) + + # 阶段 2:如果指定了配置文件,尝试读取与当前层及部署拓扑匹配的历史布局。 + # load_layer_placement 会负责记录 warning,并在任何异常或配置无效时返回 None。 + config_path = start_args.eplb_config_path + if config_path is not None: + saved_placement = load_layer_placement( + config_path, + layer_index=self.layer_index, + num_logical_experts=self.n_routed_experts, + world_size=world_size, + num_redundant_experts_per_rank=self.num_redundant_experts_per_rank, + ) + + # 阶段 3:只有完整校验通过的历史布局才会替换默认布局,使专家权重在 + # 初始化时直接加载到上一次优化后的物理槽位中。 + if saved_placement is not None: + initial_local_expert_ids_by_rank = saved_placement self.local_logics_expert_ids_list = initial_local_expert_ids_by_rank[global_rank] self.logical_to_physical_map = torch.tensor( build_logical_to_physical_map( diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 62a60875b4..786558be30 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -780,6 +780,13 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: help="""Maximum number of completed EPLB rebalances. -1 means unlimited, 0 disables dynamic rebalancing, and the default is 1.""", ) + parser.add_argument( + "--eplb_config_path", + type=str, + default=None, + help="""Path to an EPLB placement JSON file. A valid saved layout is loaded during weight + initialization, and the latest runtime layout is written back to the same path.""", + ) parser.add_argument( "--enable_fused_shared_experts", action="store_true", diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 7cb317c75b..8f7efc08c8 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -188,6 +188,7 @@ class StartArgs: enable_ep_moe: bool = field(default=False) eplb_num_redundant_experts_per_rank: int = field(default=0) eplb_rebalance_count: int = field(default=1) + eplb_config_path: Optional[str] = field(default=None) enable_fused_shared_experts: bool = field(default=False) mtp_mode: Optional[str] = field( default=None, diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index afbbc1ca49..ffcf2b0e3a 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -260,7 +260,11 @@ def init_model(self, kvargs): if self.args.eplb_num_redundant_experts_per_rank > 0: from lightllm.server.router.model_infer.mode_backend.eplb.runtime_manager import EPLBManager - self.eplb_manager = EPLBManager(self.model, self.args.eplb_rebalance_count) + self.eplb_manager = EPLBManager( + self.model, + max_rebalance_count=self.args.eplb_rebalance_count, + config_path=self.args.eplb_config_path, + ) # 启动infer_loop_thread, 启动两个线程进行推理,对于具备双batch推理折叠得场景 # 可以降低 cpu overhead,大幅提升gpu得使用率。 diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py index d63c414f14..8b68634d0e 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py @@ -5,6 +5,7 @@ from .initial import build_initial_local_expert_ids from .routing import build_logical_to_physical_map from .greedy import GreedyEPLBPlanner +from .config import load_layer_placement, save_placement_config __all__ = [ "EPLBPlanner", @@ -15,4 +16,6 @@ "LogicalToPhysicalMap", "build_initial_local_expert_ids", "build_logical_to_physical_map", + "load_layer_placement", + "save_placement_config", ] diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/config.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/config.py new file mode 100644 index 0000000000..0a7238e983 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/config.py @@ -0,0 +1,230 @@ +"""EPLB expert placement JSON loading and persistence.""" + +import json +import os +from functools import lru_cache +from typing import Any, Dict, Optional, Sequence + +from lightllm.utils.log_utils import init_logger + +from .types import ExpertPlacement, LayerPlacement + +logger = init_logger(__name__) + +EPLB_PLACEMENT_CONFIG_VERSION = 1 + + +def load_layer_placement( + config_path: str, + layer_index: int, + num_logical_experts: int, + world_size: int, + num_redundant_experts_per_rank: int, +) -> Optional[LayerPlacement]: + """读取并校验一个模型层的布局;任何错误都返回 ``None`` 以触发默认流程。""" + config = _read_config(config_path) + if config is None: + return None + + # 阶段 1:构造当前部署期望的元数据,后续各项校验和 warning 都以此为准。 + expected_metadata = { + "version": EPLB_PLACEMENT_CONFIG_VERSION, + "num_logical_experts": num_logical_experts, + "world_size": world_size, + "num_redundant_experts_per_rank": num_redundant_experts_per_rank, + } + + # 阶段 2:逐项显式校验配置元数据,便于直接定位具体的不匹配字段。 + version = config.get("version") + if version != expected_metadata["version"]: + logger.warning( + "EPLB placement config %s does not match the current deployment: version=%r, expected %r; " + "using the default initial placement", + config_path, + version, + expected_metadata["version"], + ) + return None + + config_num_logical_experts = config.get("num_logical_experts") + if config_num_logical_experts != expected_metadata["num_logical_experts"]: + logger.warning( + "EPLB placement config %s does not match the current deployment: num_logical_experts=%r, " + "expected %r; using the default initial placement", + config_path, + config_num_logical_experts, + expected_metadata["num_logical_experts"], + ) + return None + + config_world_size = config.get("world_size") + if config_world_size != expected_metadata["world_size"]: + logger.warning( + "EPLB placement config %s does not match the current deployment: world_size=%r, expected %r; " + "using the default initial placement", + config_path, + config_world_size, + expected_metadata["world_size"], + ) + return None + + config_num_redundant_experts_per_rank = config.get("num_redundant_experts_per_rank") + if config_num_redundant_experts_per_rank != expected_metadata["num_redundant_experts_per_rank"]: + logger.warning( + "EPLB placement config %s does not match the current deployment: " + "num_redundant_experts_per_rank=%r, expected %r; using the default initial placement", + config_path, + config_num_redundant_experts_per_rank, + expected_metadata["num_redundant_experts_per_rank"], + ) + return None + + # 阶段 3:读取当前模型层的物理槽布局,并校验形状、ID 范围及专家覆盖关系。 + try: + layers = config.get("layers") + assert isinstance(layers, dict), "the layers field must be a JSON object" + layer_placement = layers.get(str(layer_index)) + assert layer_placement is not None, f"layer {layer_index} is missing" + _validate_layer_placement( + layer_placement, + num_logical_experts=num_logical_experts, + world_size=world_size, + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + ) + except Exception as exc: + logger.warning( + "Layer %s in EPLB placement config %s is invalid (%s); using the default initial placement", + layer_index, + config_path, + exc, + ) + return None + + # 校验后复制一份,避免缓存中的原始 JSON 对象被运行态修改。 + return [list(rank_placement) for rank_placement in layer_placement] + + +def save_placement_config( + config_path: str, + layer_indexes: Sequence[int], + placement: ExpertPlacement, + num_logical_experts: int, + world_size: int, + num_redundant_experts_per_rank: int, +) -> bool: + """将完整 EPLB 布局原子写入配置路径,失败时仅记录 warning。""" + if len(layer_indexes) != len(placement) or len(set(layer_indexes)) != len(layer_indexes): + logger.warning( + "Failed to save EPLB placement config %s: layer indexes do not match the placement", + config_path, + ) + return False + + for layer_index, layer_placement in zip(layer_indexes, placement): + try: + _validate_layer_placement( + layer_placement, + num_logical_experts=num_logical_experts, + world_size=world_size, + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + ) + except Exception as exc: + logger.warning( + "Failed to save EPLB placement config %s: layer %s is invalid (%s)", + config_path, + layer_index, + exc, + ) + return False + + config = { + "version": EPLB_PLACEMENT_CONFIG_VERSION, + "num_logical_experts": num_logical_experts, + "world_size": world_size, + "num_redundant_experts_per_rank": num_redundant_experts_per_rank, + "layers": { + str(layer_index): [list(rank_placement) for rank_placement in layer_placement] + for layer_index, layer_placement in zip(layer_indexes, placement) + }, + } + + absolute_path = os.path.abspath(config_path) + parent_dir = os.path.dirname(absolute_path) + lock_path = f"{absolute_path}.lock" + lock_acquired = False + try: + os.makedirs(parent_dir, exist_ok=True) + # 通过 x 模式原子创建锁文件,避免多个服务进程同时写入同一个布局文件。 + with open(lock_path, "x", encoding="utf-8") as lock_file: + lock_acquired = True + lock_file.write(str(os.getpid())) + with open(absolute_path, "w", encoding="utf-8") as config_file: + json.dump(config, config_file, ensure_ascii=False, indent=2) + config_file.write("\n") + _read_config.cache_clear() + return True + except OSError as exc: + logger.warning("Failed to save EPLB placement config %s: %s", config_path, exc) + return False + finally: + if lock_acquired: + try: + os.unlink(lock_path) + except OSError as exc: + logger.warning("Failed to remove EPLB placement lock file %s: %s", lock_path, exc) + + +def _validate_layer_placement( + layer_placement: Any, + num_logical_experts: int, + world_size: int, + num_redundant_experts_per_rank: int, +) -> None: + """按顺序断言单层布局满足当前部署的全部约束。""" + # 阶段 1:校验 rank 维度和每个 rank 应持有的物理槽数量。 + assert isinstance(layer_placement, list), "the layer is not an array" + assert len(layer_placement) == world_size, f"the number of ranks is {len(layer_placement)}, expected {world_size}" + + num_physical_experts_per_rank = num_logical_experts // world_size + num_redundant_experts_per_rank + covered_experts = set() + for rank, rank_placement in enumerate(layer_placement): + # 阶段 2:依次校验每个 rank 的布局形状和 expert ID。 + assert isinstance(rank_placement, list), f"the placement for rank {rank} is not an array" + assert ( + len(rank_placement) == num_physical_experts_per_rank + ), f"rank {rank} has {len(rank_placement)} physical slots, expected {num_physical_experts_per_rank}" + assert all( + not isinstance(expert_id, bool) and isinstance(expert_id, int) for expert_id in rank_placement + ), f"rank {rank} contains a non-integer expert ID" + assert all( + 0 <= expert_id < num_logical_experts for expert_id in rank_placement + ), f"rank {rank} contains an out-of-range expert ID" + assert len(set(rank_placement)) == len(rank_placement), f"rank {rank} contains duplicate expert IDs" + covered_experts.update(rank_placement) + + # 阶段 3:确认所有 logical expert 至少存在一个物理副本。 + missing_experts = sorted(set(range(num_logical_experts)) - covered_experts) + assert not missing_experts, f"logical experts are missing: {missing_experts}" + + +@lru_cache(maxsize=None) +def _read_config(config_path: str) -> Optional[Dict[str, Any]]: + """读取并缓存配置,避免模型的每个 MoE 层重复解析同一个文件。""" + try: + with open(config_path, "r", encoding="utf-8") as config_file: + config = json.load(config_file) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + logger.warning( + "Failed to read EPLB placement config %s; using the default initial placement: %s", + config_path, + exc, + ) + return None + + if not isinstance(config, dict): + logger.warning( + "The root of EPLB placement config %s must be a JSON object; using the default initial placement", + config_path, + ) + return None + return config diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 2b24794aa4..504fd39edb 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -29,6 +29,7 @@ ExpertPlacement, GreedyEPLBPlanner, build_logical_to_physical_map, + save_placement_config, ) from .placement_plan_task import EPLBPlanTask @@ -67,7 +68,12 @@ class EPLBManager: 执行,主推理线程负责评估、轮询和提交结果。 """ - def __init__(self, model: TpPartBaseModel, max_rebalance_count: int = 1) -> None: + def __init__( + self, + model: TpPartBaseModel, + max_rebalance_count: int = 1, + config_path: Optional[str] = None, + ) -> None: # SM100 FP4 Mega-MoE 会将在线专家权重转换为独立的 kernel 布局,并使用源 tensor 的 data_ptr # 作为 key 缓存这些转换后的副本。EPLB 通过原地 copy_ 替换专家行,只改变权重内容而不会改变 # data_ptr,因此重平衡后 Mega-MoE 仍会读取旧的转换权重。在 EPLB 能够失效或更新该缓存前, @@ -80,9 +86,11 @@ def __init__(self, model: TpPartBaseModel, max_rebalance_count: int = 1) -> None # 模型与专家拓扑:初始化后保持不变。 self._weights: List[FusedMoeWeight] = weights + self.config_path = config_path self.global_rank: int = get_global_rank() self.world_size: int = get_global_world_size() assert self.world_size > 1, "EPLB requires more than one rank" + self.layer_indexes = [weight.layer_num_ for weight in weights] self._eplb_impls = [weight.fuse_moe_impl for weight in weights] first_impl = self._eplb_impls[0] @@ -331,6 +339,8 @@ def _step_transferring(self) -> None: if not transfer_batch: self.current_placement = self.target_placement elapsed = time.time() - self.rebalance_started_at + if self.global_rank == 0: + self._persist_current_placement() self._clear_route_counters() self.completed_rebalance_count += 1 del self.pending_transfer_batches @@ -451,6 +461,19 @@ def _clear_route_counters(self) -> None: for impl in self._eplb_impls: impl.route_counter.zero_() + def _persist_current_placement(self) -> None: + """由 rank 0 将当前完整布局写回启动时指定的输入/输出文件。""" + if self.config_path is None: + return + save_placement_config( + self.config_path, + layer_indexes=self.layer_indexes, + placement=self.current_placement, + num_logical_experts=self.num_logical_experts, + world_size=self.world_size, + num_redundant_experts_per_rank=self.num_redundant_experts_per_rank, + ) + def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: weights_by_id: Dict[int, FusedMoeWeight] = {} diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index c4583d9ea8..63aab46fef 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -248,6 +248,9 @@ def test_eplb_redundant_experts_default_to_disabled(): assert parser.parse_args(["--eplb_rebalance_count", "-1"]).eplb_rebalance_count == -1 assert parser.parse_args(["--eplb_rebalance_count", "0"]).eplb_rebalance_count == 0 assert StartArgs().eplb_rebalance_count == 1 + assert parser.parse_args([]).eplb_config_path is None + assert parser.parse_args(["--eplb_config_path", "/tmp/eplb.json"]).eplb_config_path == "/tmp/eplb.json" + assert StartArgs().eplb_config_path is None @pytest.mark.parametrize( @@ -823,7 +826,10 @@ def test_eplb_route_counter_has_one_entry_per_logical_expert(monkeypatch): args = type( "Args", (), - {"eplb_num_redundant_experts_per_rank": 2}, + { + "eplb_num_redundant_experts_per_rank": 2, + "eplb_config_path": None, + }, )() monkeypatch.setattr(deepgemm_module, "get_env_start_args", lambda: args) monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) @@ -1135,7 +1141,10 @@ def test_deepgemm_constructor_owns_eplb_runtime(monkeypatch): monkeypatch.setattr( deepgemm_module, "get_env_start_args", - lambda: SimpleNamespace(eplb_num_redundant_experts_per_rank=1), + lambda: SimpleNamespace( + eplb_num_redundant_experts_per_rank=1, + eplb_config_path=None, + ), ) monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) @@ -1158,6 +1167,49 @@ def cpu_zeros(*shape, **kwargs): assert not hasattr(impl, "expert_parallel_state") +def test_deepgemm_constructor_loads_saved_layout_before_weight_initialization(monkeypatch): + saved_placement = [[1, 0, 3], [2, 3, 1]] + monkeypatch.setattr( + deepgemm_module, + "get_env_start_args", + lambda: SimpleNamespace( + eplb_num_redundant_experts_per_rank=1, + eplb_config_path="/tmp/eplb.json", + ), + ) + monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) + monkeypatch.setattr( + deepgemm_module, + "load_layer_placement", + lambda path, **kwargs: ( + saved_placement + if path == "/tmp/eplb.json" + and kwargs + == { + "layer_index": 7, + "num_logical_experts": 4, + "world_size": 2, + "num_redundant_experts_per_rank": 1, + } + else None + ), + ) + original_zeros = torch.zeros + monkeypatch.setattr( + deepgemm_module.torch, + "zeros", + lambda *shape, **kwargs: original_zeros(*shape, dtype=kwargs.get("dtype")), + ) + + impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace(), layer_index=7) + + assert impl.local_logics_expert_ids_list == saved_placement[0] + expected_map = build_logical_to_physical_map(saved_placement, 4, current_rank=0) + assert impl.logical_to_physical_map.tolist() == expected_map + + def test_deepgemm_keeps_route_recording_when_rebalance_count_is_zero(monkeypatch): monkeypatch.setattr( deepgemm_module, @@ -1165,6 +1217,7 @@ def test_deepgemm_keeps_route_recording_when_rebalance_count_is_zero(monkeypatch lambda: SimpleNamespace( eplb_num_redundant_experts_per_rank=1, eplb_rebalance_count=0, + eplb_config_path=None, ), ) monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) @@ -1552,7 +1605,7 @@ def test_manager_returns_to_collecting_after_reaching_rebalance_limit(): impls = [SimpleNamespace(recording=True), SimpleNamespace(recording=True)] target_placement = [[[0, 1], [1, 0]]] manager.state = manager_module.EPLBManagerState.TRANSFERRING - manager.global_rank = 1 + manager.global_rank = 0 manager._eplb_impls = impls manager.current_placement = [[[0, 1], [0, 1]]] manager.target_placement = target_placement @@ -1560,12 +1613,15 @@ def test_manager_returns_to_collecting_after_reaching_rebalance_limit(): manager.max_rebalance_count = 1 manager.completed_rebalance_count = 0 manager._clear_route_counters = lambda: None + persisted_placements = [] + manager._persist_current_placement = lambda: persisted_placements.append(manager.current_placement) manager._step_transferring() assert manager.current_placement is target_placement assert manager.completed_rebalance_count == 1 assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert persisted_placements == [target_placement] assert all(impl.recording for impl in impls) @@ -2224,7 +2280,12 @@ def all_gather_object(output, local_expert_ids_by_layer, group): monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) logs = [] monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) - manager = manager_module.EPLBManager(type("Model", (), {})()) + monkeypatch.setattr( + manager_module, + "save_placement_config", + lambda *_args, **_kwargs: pytest.fail("manager initialization must not save the placement"), + ) + manager = manager_module.EPLBManager(type("Model", (), {})(), config_path="/tmp/eplb.json") assert not hasattr(manager, "_plan_task") assert not hasattr(manager, "pending_transfer_batches") assert manager.state is manager_module.EPLBManagerState.COLLECTING diff --git a/unit_tests/common/fused_moe/test_eplb_placement_config.py b/unit_tests/common/fused_moe/test_eplb_placement_config.py new file mode 100644 index 0000000000..509431abec --- /dev/null +++ b/unit_tests/common/fused_moe/test_eplb_placement_config.py @@ -0,0 +1,202 @@ +import json + +from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( + load_layer_placement, + save_placement_config, +) +from lightllm.server.router.model_infer.mode_backend.eplb.placement import config as config_module + + +def test_placement_config_round_trip(tmp_path): + config_path = tmp_path / "nested" / "eplb-placement.json" + placement = [ + [[0, 1, 2], [2, 3, 0]], + [[1, 0, 3], [2, 3, 1]], + ] + + assert save_placement_config( + str(config_path), + layer_indexes=[3, 7], + placement=placement, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + assert ( + load_layer_placement( + str(config_path), + layer_index=7, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + == placement[1] + ) + + saved_config = json.loads(config_path.read_text(encoding="utf-8")) + assert saved_config == { + "version": 1, + "num_logical_experts": 4, + "world_size": 2, + "num_redundant_experts_per_rank": 1, + "layers": {"3": placement[0], "7": placement[1]}, + } + assert not (config_path.parent / f"{config_path.name}.lock").exists() + + +def test_placement_config_lock_prevents_concurrent_write(tmp_path, monkeypatch): + config_path = tmp_path / "eplb-placement.json" + lock_path = tmp_path / "eplb-placement.json.lock" + config_path.write_text("original", encoding="utf-8") + lock_path.write_text("another-process", encoding="utf-8") + warnings = [] + monkeypatch.setattr(config_module.logger, "warning", lambda *args: warnings.append(args)) + + assert not save_placement_config( + str(config_path), + layer_indexes=[3], + placement=[[[0, 1, 2], [2, 3, 0]]], + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + assert config_path.read_text(encoding="utf-8") == "original" + assert lock_path.read_text(encoding="utf-8") == "another-process" + assert len(warnings) == 1 + + +def test_missing_placement_config_warns_and_falls_back(tmp_path, monkeypatch): + config_path = tmp_path / "missing.json" + warnings = [] + config_module._read_config.cache_clear() + monkeypatch.setattr(config_module.logger, "warning", lambda *args: warnings.append(args)) + + assert ( + load_layer_placement( + str(config_path), + layer_index=3, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + is None + ) + assert len(warnings) == 1 + assert "using the default initial placement" in warnings[0][0] + + +def test_malformed_json_warns_and_falls_back(tmp_path, monkeypatch): + config_path = tmp_path / "malformed.json" + config_path.write_text("{not-json", encoding="utf-8") + warnings = [] + config_module._read_config.cache_clear() + monkeypatch.setattr(config_module.logger, "warning", lambda *args: warnings.append(args)) + + assert ( + load_layer_placement( + str(config_path), + layer_index=3, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + is None + ) + assert len(warnings) == 1 + assert "using the default initial placement" in warnings[0][0] + + +def test_invalid_placement_config_warns_and_falls_back(tmp_path, monkeypatch): + config_path = tmp_path / "invalid.json" + config_path.write_text( + json.dumps( + { + "version": 1, + "num_logical_experts": 4, + "world_size": 2, + "num_redundant_experts_per_rank": 1, + "layers": {"3": [[0, 1, 1], [2, 3, 0]]}, + } + ), + encoding="utf-8", + ) + warnings = [] + config_module._read_config.cache_clear() + monkeypatch.setattr(config_module.logger, "warning", lambda *args: warnings.append(args)) + + assert ( + load_layer_placement( + str(config_path), + layer_index=3, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + is None + ) + assert len(warnings) == 1 + assert "contains duplicate expert IDs" in str(warnings[0][3]) + + +def test_missing_layer_warns_and_falls_back(tmp_path, monkeypatch): + config_path = tmp_path / "missing-layer.json" + config_path.write_text( + json.dumps( + { + "version": 1, + "num_logical_experts": 4, + "world_size": 2, + "num_redundant_experts_per_rank": 1, + "layers": {}, + } + ), + encoding="utf-8", + ) + warnings = [] + config_module._read_config.cache_clear() + monkeypatch.setattr(config_module.logger, "warning", lambda *args: warnings.append(args)) + + assert ( + load_layer_placement( + str(config_path), + layer_index=3, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + is None + ) + assert len(warnings) == 1 + assert "is missing" in str(warnings[0][3]) + + +def test_topology_mismatch_warns_and_falls_back(tmp_path, monkeypatch): + config_path = tmp_path / "mismatch.json" + config_path.write_text( + json.dumps( + { + "version": 1, + "num_logical_experts": 8, + "world_size": 2, + "num_redundant_experts_per_rank": 1, + "layers": {}, + } + ), + encoding="utf-8", + ) + warnings = [] + config_module._read_config.cache_clear() + monkeypatch.setattr(config_module.logger, "warning", lambda *args: warnings.append(args)) + + assert ( + load_layer_placement( + str(config_path), + layer_index=3, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + ) + is None + ) + assert len(warnings) == 1 + assert "does not match the current deployment" in warnings[0][0] From 1b1bdc59aedf9bbccf1beb2527b2d830f51a30ba Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 07:37:47 +0000 Subject: [PATCH 58/72] refactor(eplb): carry source slot in transfer info --- .../mode_backend/eplb/expert_transfer.py | 60 +++++++++++-------- unit_tests/common/fused_moe/test_eplb.py | 60 ++++++++++--------- 2 files changed, 66 insertions(+), 54 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py index cc452cb5ad..67e0c8e901 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py @@ -31,14 +31,15 @@ class ExpertTensorBuffer: class EPLBTransferInfo: """单个逻辑专家的一次传输描述。 - ``layer_index`` 是专家权重在 EPLB 层列表中的下标。源 rank 使用 - ``source_logical_expert_id`` 定位当前本地物理行;目标 rank 将收到的 - pinned memory 数据写入 ``dest_local_expert_index`` 指定的本地物理行。 + ``layer_index`` 和 ``source_logical_expert_id`` 标识需要传输的专家; + ``source_rank``、``source_local_expert_index`` 描述当前物理槽, + ``dest_rank``、``dest_local_expert_index`` 描述目标物理槽。 """ - source_rank: int layer_index: int source_logical_expert_id: int + source_rank: int + source_local_expert_index: int dest_rank: int dest_local_expert_index: int @@ -65,9 +66,9 @@ class PinnedMemoryEPLBTransfer: ``status`` 会变为 :attr:`TransferStatus.SUCCEEDED`,收到的数据保存在 ``tensor_buffers``。EPLBManager 在主循环的安全边界同步提交这些数据。 - 每个对象只表示构造函数中 ``transfer_info`` 指定的一次传输。逻辑专家 ID - 在源 rank 上通过该层当前的本地专家列表解析为物理行;目标物理槽位不属于 - 传输职责,由 manager 根据目标 placement 决定。 + 每个对象只表示构造函数中 ``transfer_info`` 指定的一次传输。源 rank 直接 + 读取 ``source_local_expert_index`` 指定的物理行;目标物理槽位不属于传输 + 职责,由 manager 根据目标 placement 决定。 """ def __init__( @@ -87,7 +88,6 @@ def __init__( # (或非量化权重),以及配套的 weight_scale、weight_zero_point 等量化 # 信息。后续会为每项张量创建对应的 pinned row,确保专家状态完整迁移。 named_live_tensors: List[NamedTensor] = extract_eplb_expert_tensors(layer_weight) - self._local_logical_expert_ids: List[int] = layer_weight.fuse_moe_impl.local_logics_expert_ids_list self._device: torch.device = named_live_tensors[0][1].device # 只有源和目标 rank 需要保存该专家的 pinned row。源 rank 用它作为 @@ -134,13 +134,10 @@ def _run_transfer(self) -> None: transfer_info: EPLBTransferInfo = self.transfer_info if self._is_source_rank: torch.cuda.set_device(self._device) - source_local_expert_index: int = self._local_logical_expert_ids.index( - transfer_info.source_logical_expert_id - ) with torch.cuda.stream(self._device_to_host_stream): for tensor_buffer in self.tensor_buffers: tensor_buffer.pinned_row.copy_( - tensor_buffer.live_tensor[source_local_expert_index], + tensor_buffer.live_tensor[transfer_info.source_local_expert_index], non_blocking=True, ) # Gloo 读取 pinned row 前,源 rank 必须等待 GPU -> CPU 拷贝完成。 @@ -175,14 +172,15 @@ def _build_p2p_message_tag(self, tensor_name: str) -> int: Python ``hash`` 会因进程随机种子不同而产生不同结果,因此这里使用稳定的 CRC32,并限制到 Gloo 可安全使用的有符号 31 位整数范围。标识中包含层、 - 源 rank、目标 rank、逻辑专家和张量名称,避免依赖张量列表的隐式顺序。 + 逻辑专家、源物理槽、目标物理槽和张量名称,避免依赖张量列表的隐式顺序。 """ transfer_info: EPLBTransferInfo = self.transfer_info message_identity = ( f"{transfer_info.layer_index}:" + f"{transfer_info.source_logical_expert_id}:" f"{transfer_info.source_rank}:" + f"{transfer_info.source_local_expert_index}:" f"{transfer_info.dest_rank}:" - f"{transfer_info.source_logical_expert_id}:" f"{transfer_info.dest_local_expert_index}:" f"{tensor_name}" ) @@ -281,10 +279,9 @@ def build_transfer_plan( # # 稳定副本不会出现在任何任务的目标位置,因此可以反复读取而没有覆盖 # 风险。只有不存在稳定副本时,才循环使用该专家当前已有的所有副本。 - # pending_transfers 的每一项为 (source_slot, transfer_info)。source_slot - # 只用于规划覆盖依赖,真正执行任务所需的信息保存在 transfer_info 中。 + # 源槽位完整保存在 transfer_info 中,后续依赖分析和实际传输共用同一份信息。 source_use_count = [0] * num_logical_experts - pending_transfers: List[tuple[Slot, EPLBTransferInfo]] = [] + pending_transfers: List[EPLBTransferInfo] = [] for destination_rank, (current_row, target_row) in enumerate(zip(current_placement, target_placement)): for destination_local_expert_index in range(num_local_experts_per_rank): current_expert_id = current_row[destination_local_expert_index] @@ -295,18 +292,21 @@ def build_transfer_plan( source_slot = source_slots[source_use_count[target_expert_id] % len(source_slots)] source_use_count[target_expert_id] += 1 transfer_info = EPLBTransferInfo( - source_rank=source_slot[0], layer_index=layer_index, source_logical_expert_id=target_expert_id, + source_rank=source_slot[0], + source_local_expert_index=source_slot[1], dest_rank=destination_rank, dest_local_expert_index=destination_local_expert_index, ) - pending_transfers.append((source_slot, transfer_info)) + pending_transfers.append(transfer_info) # 阶段 3:按照槽位覆盖依赖,将任务拆成可安全提交的执行批次。 transfer_batches: List[List[EPLBTransferInfo]] = [] while pending_transfers: - source_slots = {source_slot for source_slot, _ in pending_transfers} + source_slots = { + (transfer_info.source_rank, transfer_info.source_local_expert_index) for transfer_info in pending_transfers + } # 3.1 收集当前拓扑层次的全部安全任务。source_slots 是当前仍需保护的 # 槽位集合:只要某个槽位中的专家尚未完成最后一次读取,该槽位就仍在 @@ -320,13 +320,13 @@ def build_transfer_plan( # 这里不能在找到第一个任务后立即修改 source_slots。只有整批提交并从 # pending 中移除后,下一层目标槽位才真正变得安全。 safe_transfer_batch: List[EPLBTransferInfo] = [] - remaining_transfers: List[tuple[Slot, EPLBTransferInfo]] = [] - for source_slot, transfer_info in pending_transfers: + remaining_transfers: List[EPLBTransferInfo] = [] + for transfer_info in pending_transfers: destination_slot = (transfer_info.dest_rank, transfer_info.dest_local_expert_index) if destination_slot not in source_slots: safe_transfer_batch.append(transfer_info) else: - remaining_transfers.append((source_slot, transfer_info)) + remaining_transfers.append(transfer_info) if safe_transfer_batch: # 安全任务之间没有原子提交要求,但若同一 rank 在一个批次中参与 @@ -368,9 +368,13 @@ def build_transfer_plan( # B 不在源集合中,所以 ``S -> B`` 会先作为安全任务移除;剩余的 # ``S -> A、A -> S`` 才会进入这里,并且每个源都只对应一个目标。 assert len(source_slots) == len(pending_transfers) - transfer_by_source_slot = dict(pending_transfers) + transfer_by_source_slot = { + (transfer_info.source_rank, transfer_info.source_local_expert_index): transfer_info + for transfer_info in pending_transfers + } - cycle_start_slot = pending_transfers[0][0] + first_transfer = pending_transfers[0] + cycle_start_slot = (first_transfer.source_rank, first_transfer.source_local_expert_index) source_slot = cycle_start_slot cycle_batch: List[EPLBTransferInfo] = [] cycle_source_slots: set[Slot] = set() @@ -388,6 +392,10 @@ def build_transfer_plan( source_slot = destination_slot transfer_batches.append(cycle_batch) - pending_transfers = [transfer for transfer in pending_transfers if transfer[0] not in cycle_source_slots] + pending_transfers = [ + transfer_info + for transfer_info in pending_transfers + if (transfer_info.source_rank, transfer_info.source_local_expert_index) not in cycle_source_slots + ] return transfer_batches diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 63aab46fef..020d3c031f 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -654,6 +654,10 @@ def test_transfer_plan_respects_explicit_target_slots(): transfer_infos = [transfer_info for transfer_batch in plan for transfer_info in transfer_batch] assert all(info.layer_index == 3 for info in transfer_infos) + assert all( + current[info.source_rank][info.source_local_expert_index] == info.source_logical_expert_id + for info in transfer_infos + ) assert {(info.dest_rank, info.source_logical_expert_id) for info in transfer_infos} == { (rank, target[rank][slot]) for rank in range(4) for slot in range(2, 4) } @@ -737,8 +741,8 @@ def test_transfer_planner_combines_all_layer_batches(monkeypatch): current_placement = [[[0, 1], [2, 3]], [[0, 2], [1, 3]]] target_placement = [[[2, 1], [0, 3]], [[0, 3], [1, 2]]] transfer_infos = [ - EPLBTransferInfo(1, 0, 2, 0, 0), - EPLBTransferInfo(1, 1, 3, 0, 1), + EPLBTransferInfo(0, 2, 1, 0, 0, 0), + EPLBTransferInfo(1, 3, 1, 1, 0, 1), ] calls = [] @@ -1337,8 +1341,8 @@ def test_transfer_plan_uses_stable_current_expert_source(): plan = build_transfer_plan(current, target, 5, num_logical_experts=8, world_size=4) assert plan == [ [ - EPLBTransferInfo(1, 5, 6, 0, 2), - EPLBTransferInfo(2, 5, 4, 2, 3), + EPLBTransferInfo(5, 6, 1, 2, 0, 2), + EPLBTransferInfo(5, 4, 2, 0, 2, 3), ], ] @@ -1350,8 +1354,8 @@ def test_transfer_plan_reuses_stable_source_for_repeated_expert(): second = build_transfer_plan(current, target, 5, 8, 4) assert first == second assert first == [ - [EPLBTransferInfo(2, 5, 4, 0, 2)], - [EPLBTransferInfo(2, 5, 4, 0, 3)], + [EPLBTransferInfo(5, 4, 2, 0, 0, 2)], + [EPLBTransferInfo(5, 4, 2, 2, 0, 3)], ] @@ -1363,8 +1367,8 @@ def test_transfer_plan_keeps_primary_slot_swap_in_one_atomic_batch(): assert plan == [ [ - EPLBTransferInfo(1, 0, 2, 0, 0), - EPLBTransferInfo(0, 0, 0, 1, 0), + EPLBTransferInfo(0, 2, 1, 0, 0, 0), + EPLBTransferInfo(0, 0, 0, 0, 1, 0), ] ] @@ -1377,24 +1381,26 @@ def test_transfer_plan_keeps_three_way_cycle_in_one_atomic_batch(): assert plan == [ [ - EPLBTransferInfo(1, 0, 1, 0, 0), - EPLBTransferInfo(0, 0, 0, 2, 0), - EPLBTransferInfo(2, 0, 2, 1, 0), + EPLBTransferInfo(0, 1, 1, 0, 0, 0), + EPLBTransferInfo(0, 0, 0, 0, 2, 0), + EPLBTransferInfo(0, 2, 2, 0, 1, 0), ] ] def test_p2p_message_tag_is_stable_and_identifies_transfer_tensor(): transfer = object.__new__(PinnedMemoryEPLBTransfer) - transfer.transfer_info = EPLBTransferInfo(1, 5, 4, 0, 2) + transfer.transfer_info = EPLBTransferInfo(5, 4, 1, 0, 0, 2) weight_tag = transfer._build_p2p_message_tag("w13.weight") assert weight_tag == transfer._build_p2p_message_tag("w13.weight") assert 0 <= weight_tag <= 0x7FFFFFFF assert weight_tag != transfer._build_p2p_message_tag("w13.weight_scale") - transfer.transfer_info = EPLBTransferInfo(1, 5, 6, 0, 2) + transfer.transfer_info = EPLBTransferInfo(5, 6, 1, 0, 0, 2) + assert weight_tag != transfer._build_p2p_message_tag("w13.weight") + transfer.transfer_info = EPLBTransferInfo(5, 4, 1, 1, 0, 2) assert weight_tag != transfer._build_p2p_message_tag("w13.weight") - transfer.transfer_info = EPLBTransferInfo(1, 5, 4, 0, 3) + transfer.transfer_info = EPLBTransferInfo(5, 4, 1, 0, 0, 3) assert weight_tag != transfer._build_p2p_message_tag("w13.weight") @@ -1453,11 +1459,11 @@ def record_copy(tensor, source, non_blocking=False): ] transfers = [ SimpleNamespace( - transfer_info=EPLBTransferInfo(1, 0, 4, 0, 3), + transfer_info=EPLBTransferInfo(0, 4, 1, 1, 0, 3), tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -4))], ), SimpleNamespace( - transfer_info=EPLBTransferInfo(1, 0, 5, 0, 4), + transfer_info=EPLBTransferInfo(0, 5, 1, 2, 0, 4), tensor_buffers=[ExpertTensorBuffer("weight", live, torch.full((4,), -5))], ), ] @@ -1489,9 +1495,9 @@ def start(self): def is_finished(self): return self.finished - remote_info = EPLBTransferInfo(0, 0, 2, 2, 2) - local_info0 = EPLBTransferInfo(1, 0, 3, 3, 2) - local_info1 = EPLBTransferInfo(0, 1, 4, 1, 2) + remote_info = EPLBTransferInfo(0, 2, 0, 0, 2, 2) + local_info0 = EPLBTransferInfo(0, 3, 1, 1, 3, 2) + local_info1 = EPLBTransferInfo(1, 4, 0, 0, 1, 2) finished_by_info = {local_info0: False, local_info1: True} starts = [] manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) @@ -1668,7 +1674,7 @@ def is_finished(self): original_overlap_stream = g_infer_context.overlap_stream manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - transfer_info = EPLBTransferInfo(0, 0, 0, 0, 0) + transfer_info = EPLBTransferInfo(0, 0, 0, 0, 0, 0) transfer = Transfer(live, received, transfer_info) manager.active_transfers = [transfer] manager.active_transfer_batch = [transfer_info] @@ -1800,8 +1806,8 @@ def test_manager_plans_transfers_asynchronously_before_entering_transferring(mon ) transfer_infos = [ - EPLBTransferInfo(0, 0, 2, 1, 2), - EPLBTransferInfo(0, 1, 0, 1, 2), + EPLBTransferInfo(0, 2, 0, 0, 1, 2), + EPLBTransferInfo(1, 0, 0, 0, 1, 2), ] transfer_planners = [] @@ -2048,9 +2054,8 @@ def synchronize(self): transfer._is_source_rank = True transfer._is_destination_rank = False transfer._p2p_group = object() - transfer.transfer_info = EPLBTransferInfo(0, 0, 5, 1, 2) + transfer.transfer_info = EPLBTransferInfo(0, 5, 0, 1, 1, 2) transfer._device_to_host_stream = Stream() - transfer._local_logical_expert_ids = [4, 5] transfer.tensor_buffers = [ ExpertTensorBuffer( "weight", @@ -2089,9 +2094,8 @@ def synchronize(self): transfer._is_source_rank = True transfer._is_destination_rank = True transfer._p2p_group = object() - transfer.transfer_info = EPLBTransferInfo(0, 0, 5, 0, 1) + transfer.transfer_info = EPLBTransferInfo(0, 5, 0, 0, 0, 1) transfer._device_to_host_stream = Stream() - transfer._local_logical_expert_ids = [5] transfer.tensor_buffers = [ ExpertTensorBuffer( "weight", @@ -2118,7 +2122,7 @@ def test_pinned_transfer_exits_process_on_failure(monkeypatch): transfer._device = "cuda:0" transfer._is_source_rank = False transfer._is_destination_rank = True - transfer.transfer_info = EPLBTransferInfo(1, 0, 3, 0, 2) + transfer.transfer_info = EPLBTransferInfo(0, 3, 1, 0, 0, 2) transfer._p2p_group = object() transfer.tensor_buffers = [ ExpertTensorBuffer( @@ -2155,7 +2159,7 @@ def synchronize(self): transfer._device = "cuda:0" transfer._is_source_rank = False transfer._is_destination_rank = True - transfer.transfer_info = EPLBTransferInfo(1, 0, 3, 0, 2) + transfer.transfer_info = EPLBTransferInfo(0, 3, 1, 0, 0, 2) transfer._p2p_group = object() transfer._device_to_host_stream = Stream() transfer.tensor_buffers = [ From 942e0b59f61f947a5c2d480b7f6e1ad6064b613e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 09:11:49 +0000 Subject: [PATCH 59/72] feat(eplb): add topology-aware expert dispatch modes --- .../fused_moe/impl/deepgemm_impl.py | 3 + .../triton_kernel/fused_moe/eplb_topk_ids.py | 79 ++++-- .../mode_backend/eplb/placement/routing.py | 96 +++++--- .../mode_backend/eplb/runtime_manager.py | 3 + unit_tests/common/fused_moe/test_eplb.py | 229 ++++++++++++++---- 5 files changed, 314 insertions(+), 96 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 add55795ac..c64002ed7b 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 @@ -16,6 +16,7 @@ from lightllm.utils.dist_utils import ( get_global_rank, get_global_world_size, + get_node_world_size, ) from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( fused_experts, @@ -82,6 +83,7 @@ def _init_eplb_runtime(self): initial_local_expert_ids_by_rank, self.n_routed_experts, current_rank=global_rank, + node_world_size=get_node_world_size(), ), dtype=torch.int32, ).cuda() @@ -147,6 +149,7 @@ def _prepare_expert_execution( logical_to_physical_map=self.logical_to_physical_map, logical_expert_counter=self.route_counter, update_logical_expert_counter=self.recording, + mode="current_gpu_first", ) return topk_weights, topk_ids diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py index a6fcb7c63e..c26d7eadc3 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py @@ -19,6 +19,7 @@ def _eplb_repair_topk_ids_kernel( logical_to_physical_map_ptr, logical_to_physical_map_row_stride, logical_expert_counter_ptr, + DISPATCH_MODE: tl.constexpr, UPDATE_LOGICAL_EXPERT_COUNTER: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): @@ -38,24 +39,51 @@ def _eplb_repair_topk_ids_kernel( sem="relaxed", ) - # 阶段 3:定位每个 logical expert 的打包映射行。第 0 列保存有效 - # physical 副本数,第 1 列标记当前 rank 是否有本地副本,第 2 列起 - # 保存可参与路由的 physical expert ID;本地副本存在时固定放在第一个槽。 + # 阶段 3:定位每个 logical expert 的打包映射行。固定头部的布局为: + # + # [0] 所有 rank 上的有效副本总数 + # [1] 当前节点上的有效副本数,包含本卡副本 + # [2] 当前 GPU 上的有效副本数 + # [3:] 按本卡、本节点其他卡、其他节点排列的 physical expert IDs + # + # 因为三层候选在 physical ID 列表中都是连续前缀,所以选择对应层级的 + # count 后,可以直接对 [3:] 的前 count 项做 hash。 map_row_offsets = logical_expert_ids * logical_to_physical_map_row_stride - num_valid_replicas = tl.load(logical_to_physical_map_ptr + map_row_offsets, mask=valid_mask, other=1) - has_local_replica = tl.load(logical_to_physical_map_ptr + map_row_offsets + 1, mask=valid_mask, other=0) + num_global_replicas = tl.load(logical_to_physical_map_ptr + map_row_offsets, mask=valid_mask, other=1) + num_node_replicas = tl.load(logical_to_physical_map_ptr + map_row_offsets + 1, mask=valid_mask, other=0) + num_current_gpu_replicas = tl.load( + logical_to_physical_map_ptr + map_row_offsets + 2, + mask=valid_mask, + other=0, + ) - # 阶段 4:若当前 rank 持有该专家,强制选择第一个槽位以避免跨 rank - # 通信;否则用 token 下标和 logical expert ID 生成稳定 hash,在有效 - # 副本范围内选择槽位,避免固定 token 位置长期偏向同一个副本。 + # 阶段 4:根据调用方显式指定的分发模式选择参与 hash 的候选前缀。 + if DISPATCH_MODE == 0: + # current_gpu_first: 本卡 -> 全局。 + # TODO: 等 EPLB 布局算法支持节点拓扑感知后,再考虑增加 + # 本卡 -> 本节点 -> 全局的分层回退行为。 + num_preferred_replicas = tl.where( + num_current_gpu_replicas > 0, + num_current_gpu_replicas, + num_global_replicas, + ) + elif DISPATCH_MODE == 1: + # current_node_first: 本节点 -> 全局;不单独优先本卡。 + num_preferred_replicas = tl.where( + num_node_replicas > 0, + num_node_replicas, + num_global_replicas, + ) + else: + # global_first: 直接在所有 rank 的有效副本间分发。 + num_preferred_replicas = num_global_replicas token_indices = topk_id_offsets // top_k - hashed_replica_indices = _replica_index(token_indices, logical_expert_ids, num_valid_replicas) - selected_replica_indices = tl.where(has_local_replica != 0, 0, hashed_replica_indices) + selected_replica_indices = _replica_index(token_indices, logical_expert_ids, num_preferred_replicas) # 阶段 5:读取选中槽位的 physical expert ID 并写入新的输出 tensor。 # logical_topk_ids 只读,后续 callback 仍可安全观察原始逻辑路由结果。 physical_expert_ids = tl.load( - logical_to_physical_map_ptr + map_row_offsets + 2 + selected_replica_indices, + logical_to_physical_map_ptr + map_row_offsets + 3 + selected_replica_indices, mask=valid_mask, other=-1, ) @@ -68,28 +96,48 @@ def eplb_repair_topk_ids( logical_to_physical_map: torch.Tensor, logical_expert_counter: torch.Tensor, update_logical_expert_counter: bool, + mode: str, ) -> torch.Tensor: """将 logical top-k ID 转换为当前 EPLB 布局中的 physical expert ID。 参数: logical_topk_ids: logical expert ID,shape 为 ``[num_tokens, top_k]``。 logical_to_physical_map: 打包路由表,shape 为 - ``[num_logical_experts, 2 + routing_slots]``。第 0 列保存有效副本 - 数量,第 1 列标记当前 rank 是否有本地副本,其余列保存 physical - expert ID;存在本地副本时,其 physical ID 固定放在第一个槽位。 + ``[num_logical_experts, 3 + routing_slots]``,单行布局为: + + ``[global_count, node_count, current_gpu_count, physical_ids..., padding...]`` + + 三个计数依次表示全局、当前节点和当前 GPU 上的有效副本数量。 + physical IDs 按本卡、本节点其他卡、其他节点排列;具体参与分发的 + 候选前缀由 ``mode`` 决定。 有效副本之后未使用的 padding 槽位为 -1,kernel 不会读取它们。 logical_expert_counter: 每个 logical expert 的累计路由次数,shape 为 ``[num_logical_experts]``。 update_logical_expert_counter: 是否将本次 logical 路由结果累计到 ``logical_expert_counter``。固定布局不需要动态重排时可以关闭。 + mode: 必须显式指定的副本分发模式,不提供默认值: + + * ``current_gpu_first``:本卡优先,没有本卡副本时回退到全局; + * ``current_node_first``:本节点优先,没有节点内副本时回退到全局; + * ``global_first``:直接在全局全部有效副本间分发。 返回: physical expert ID,shape 为 ``[num_tokens, top_k]``。 """ + dispatch_mode_ids = { + "current_gpu_first": 0, + "current_node_first": 1, + "global_first": 2, + } + assert ( + mode in dispatch_mode_ids + ), f"unsupported EPLB dispatch mode {mode!r}; expected one of {tuple(dispatch_mode_ids)}" + dispatch_mode = dispatch_mode_ids[mode] + assert logical_topk_ids.is_contiguous() assert logical_topk_ids.ndim == 2 assert logical_to_physical_map.ndim == 2 - assert logical_to_physical_map.shape[1] > 2 + assert logical_to_physical_map.shape[1] > 3 assert logical_to_physical_map.stride(1) == 1 assert logical_expert_counter.ndim == 1 assert logical_expert_counter.shape[0] == logical_to_physical_map.shape[0] @@ -106,6 +154,7 @@ def eplb_repair_topk_ids( logical_to_physical_map_ptr=logical_to_physical_map, logical_to_physical_map_row_stride=logical_to_physical_map.stride(0), logical_expert_counter_ptr=logical_expert_counter, + DISPATCH_MODE=dispatch_mode, UPDATE_LOGICAL_EXPERT_COUNTER=update_logical_expert_counter, BLOCK_SIZE=block_size, num_warps=4, diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/routing.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/routing.py index d03ce297cb..9fcb69c679 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/routing.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/routing.py @@ -13,6 +13,7 @@ def build_logical_to_physical_map( rank_to_logic_expert_ids: LayerPlacement, num_logical_experts: int, current_rank: int, + node_world_size: int, ) -> LogicalToPhysicalMap: """使用普通 CPU list 构建单层 logical 到 physical expert 的路由表。 @@ -20,11 +21,19 @@ def build_logical_to_physical_map( ``[num_ranks, num_physical_experts_per_rank]``,每行包含该 rank 的全部 物理专家。 - 返回值的 shape 为 ``[num_logical_experts, 2 + routing_slots]``。每一行 - 对应一个 logical expert:第 0 项是有效副本数,第 1 项 - 标记 ``current_rank`` 是否持有本地副本,第 2 项起是 physical expert ID。 - 如果本 rank 持有副本,该副本固定放在第一个路由槽;有效副本之后未使用 - 的固定宽度 padding 槽位填充为 ``-1``。 + 返回值的 shape 为 ``[num_logical_experts, 3 + routing_slots]``。每一行的 + 可视化结构如下: + + ``[global_count, node_count, current_gpu_count, physical_ids..., -1 padding...]`` + + * ``global_count``:所有 rank 上的有效副本总数; + * ``node_count``:当前节点上的有效副本数,包含本卡副本; + * ``current_gpu_count``:当前 GPU 上的有效副本数; + * ``physical_ids``:依次按本卡、本节点其他卡、其他节点排列的副本 ID。 + + 路由时优先使用最靠近当前 GPU 的非空候选集合:先使用本卡副本,其次使用 + 本节点副本,当前节点没有副本时才使用所有 rank 的副本。有效副本之后未 + 使用的固定宽度槽位填充为 ``-1``。 本函数只负责 CPU 元数据计算。调用方需要设备 Tensor 时,应在函数外 显式执行 ``torch.tensor(...)``。 @@ -38,10 +47,12 @@ def build_logical_to_physical_map( num_primary_experts_per_rank = num_logical_experts // num_ranks num_redundant_experts_per_rank = num_physical_experts_per_rank - num_primary_experts_per_rank assert num_redundant_experts_per_rank >= 0 - # 阶段 2:计算固定路由槽宽度。该宽度沿用初始化时“一个基础副本加上 - # 全部冗余槽”的容量上界;动态布局不再要求基础副本位于固定槽位。 - num_routing_slots = 1 + num_ranks * num_redundant_experts_per_rank + # 阶段 2:使用整个 world 的物理槽位总数作为固定路由槽宽度。实际候选 + # 仍只写入有效副本,其余槽位统一 padding 为 -1。 + num_routing_slots = num_ranks * num_physical_experts_per_rank assert 0 <= current_rank < num_ranks + assert 0 < node_world_size <= num_ranks + assert num_ranks % node_world_size == 0 # 阶段 3:把“物理槽 -> logical expert”的完整布局反转为 # “logical expert -> 全部物理槽”,得到每个专家的候选副本列表。 @@ -50,25 +61,33 @@ def build_logical_to_physical_map( num_logical_experts, ) - # 阶段 4:对每个候选列表做稳定排序。本 rank 的 physical ID 排在前面, - # 因而后续只需查看第一个候选,就能判断和选择本地副本。 + # 阶段 4:对每个候选列表做稳定排序。当前 rank 的 physical ID 排在最前, + # 同节点其他 rank 次之,跨节点副本最后。 _sort_physical_ids_by_locality( physical_ids_by_logical_expert, current_rank, num_physical_experts_per_rank, + node_world_size, ) - local_physical_id_start = current_rank * num_physical_experts_per_rank - local_physical_id_end = local_physical_id_start + num_physical_experts_per_rank - # 阶段 5:逐个 logical expert 打包固定宽度的路由行。实际副本不足固定 # 宽度时,剩余槽位使用 -1 padding;kernel 只会索引有效副本范围。 logical_to_physical_map = [] + current_node = current_rank // node_world_size + current_node_rank_start = current_node * node_world_size + current_node_rank_end = current_node_rank_start + node_world_size for physical_expert_ids in physical_ids_by_logical_expert: - has_local_replica = local_physical_id_start <= physical_expert_ids[0] < local_physical_id_end + replica_ranks = [ + physical_expert_id // num_physical_experts_per_rank for physical_expert_id in physical_expert_ids + ] + num_node_replicas = sum( + current_node_rank_start <= replica_rank < current_node_rank_end for replica_rank in replica_ranks + ) + num_current_gpu_replicas = sum(replica_rank == current_rank for replica_rank in replica_ranks) logical_to_physical_map.append( _build_routing_row( physical_expert_ids=physical_expert_ids, - has_local_replica=has_local_replica, + num_node_replicas=num_node_replicas, + num_current_gpu_replicas=num_current_gpu_replicas, num_routing_slots=num_routing_slots, ) ) @@ -100,46 +119,53 @@ def _sort_physical_ids_by_locality( physical_ids_by_logical_expert: list[list[int]], current_rank: int, num_physical_experts_per_rank: int, + node_world_size: int, ) -> None: - """按照 physical ID 是否属于当前 rank,对每个副本列表稳定排序。 + """按当前 rank、当前节点、其他节点的优先级稳定排序副本。 - 本地 physical ID 的排序键为 0,其他 physical ID 的排序键为 1。因此 - 当前 rank 持有的副本会移动到列表前面,同时本地副本之间、远端副本 - 之间的原始顺序保持不变。当前 rank 没有副本的列表顺序不会发生变化。 + 当前 rank 的排序键为 0,同节点其他 rank 为 1,其他节点为 2。同一优先级 + 内保持原 physical ID 顺序不变。 """ - local_physical_id_start = current_rank * num_physical_experts_per_rank - local_physical_id_end = local_physical_id_start + num_physical_experts_per_rank + current_node = current_rank // node_world_size + + def locality_priority(physical_expert_id: int) -> int: + physical_rank = physical_expert_id // num_physical_experts_per_rank + if physical_rank == current_rank: + return 0 + if physical_rank // node_world_size == current_node: + return 1 + return 2 for physical_expert_ids in physical_ids_by_logical_expert: # list.sort 是稳定排序:排序键相同时,physical ID 的原始顺序不变。 - physical_expert_ids.sort( - key=lambda physical_expert_id: ( - 0 if local_physical_id_start <= physical_expert_id < local_physical_id_end else 1 - ) - ) + physical_expert_ids.sort(key=locality_priority) def _build_routing_row( physical_expert_ids: list[int], - has_local_replica: bool, + num_node_replicas: int, + num_current_gpu_replicas: int, num_routing_slots: int, ) -> list[int]: """将一个 logical expert 的候选 physical IDs 打包为固定宽度路由行。 - ``physical_expert_ids`` 已由调用方完成本地优先的稳定排序,所以本函数 + ``physical_expert_ids`` 已由调用方完成拓扑优先的稳定排序,所以本函数 不再依赖 ``current_rank``。列表长度就是该 logical expert 的有效物理 副本数,无需额外传入容易失配的副本数量。 """ # 阶段 1:候选列表包含该专家的全部物理副本,其长度就是有效副本数。 - num_valid_replicas = len(physical_expert_ids) - assert 0 < num_valid_replicas <= num_routing_slots + num_global_replicas = len(physical_expert_ids) + assert 0 < num_global_replicas <= num_routing_slots + assert 0 <= num_current_gpu_replicas <= num_node_replicas <= num_global_replicas # 阶段 2:有效槽位直接保存稳定排序后的候选;固定宽度中未使用的尾部 - # 槽位统一填充 -1。kernel 的副本索引严格小于 num_valid_replicas, + # 槽位统一填充 -1。kernel 的副本索引严格小于 num_global_replicas, # 因而不会读取 padding。 - num_padding_slots = num_routing_slots - num_valid_replicas + num_padding_slots = num_routing_slots - num_global_replicas routing_slots = physical_expert_ids + [-1] * num_padding_slots - # 阶段 3:第 0 列保存 kernel 参与 hash 的有效副本数;第 1 列标记是否 - # 存在本地副本;后续列保存按本地优先顺序排列的 physical IDs 和 -1 padding。 - return [num_valid_replicas, int(has_local_replica), *routing_slots] + # 阶段 3:将三层有效副本计数放在固定头部,后面拼接按拓扑优先级排序的 + # physical IDs 和 -1 padding: + # + # [global_count, node_count, current_gpu_count, physical_ids..., -1 padding...] + return [num_global_replicas, num_node_replicas, num_current_gpu_replicas, *routing_slots] diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 504fd39edb..ac0f531b5b 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -13,6 +13,7 @@ from lightllm.utils.dist_utils import ( get_global_rank, get_global_world_size, + get_node_world_size, ) from lightllm.utils.device_utils import is_sm100_gpu from lightllm.utils.envs_utils import get_eplb_step_interval @@ -90,6 +91,7 @@ def __init__( self.global_rank: int = get_global_rank() self.world_size: int = get_global_world_size() assert self.world_size > 1, "EPLB requires more than one rank" + self.node_world_size: int = get_node_world_size() self.layer_indexes = [weight.layer_num_ for weight in weights] self._eplb_impls = [weight.fuse_moe_impl for weight in weights] @@ -436,6 +438,7 @@ def _publish_layer_metadata(self, layer_index: int) -> None: self.current_placement[layer_index], self.num_logical_experts, current_rank=self.global_rank, + node_world_size=self.node_world_size, ), dtype=torch.int32, pin_memory=True, diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 020d3c031f..e9c55637d3 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -64,8 +64,8 @@ def _test_moe_impl( if eplb: if route_counter is None: route_counter = torch.zeros((num_logical_experts,), dtype=torch.int64) - logical_to_physical_map = torch.zeros((num_logical_experts, world_size + 2), dtype=torch.int32) - logical_to_physical_map[:, 0] = 1 + logical_to_physical_map = torch.zeros((num_logical_experts, world_size + 3), dtype=torch.int32) + logical_to_physical_map[:, :3] = 1 else: num_redundant_experts_per_rank = 0 return SimpleNamespace( @@ -561,23 +561,34 @@ def load_zero_point(expert, local, _weights): def test_logical_to_physical_map_selects_one_physical_expert(): rank_to_logic_expert_ids = [[0, 1, 2, 3], [2, 3, 0, 1]] - logical_to_physical = build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=4, current_rank=0) + logical_to_physical = build_logical_to_physical_map( + rank_to_logic_expert_ids, + num_logical_experts=4, + current_rank=0, + node_world_size=2, + ) assert isinstance(logical_to_physical, list) assert len(logical_to_physical) == 4 - assert all(len(row) == 7 for row in logical_to_physical) + assert all(len(row) == 11 for row in logical_to_physical) assert [row[0] for row in logical_to_physical] == [2, 2, 2, 2] - assert [row[1] for row in logical_to_physical] == [1, 1, 1, 1] - assert [row[2] for row in logical_to_physical] == [0, 1, 2, 3] - assert all(physical_id >= 0 for row in logical_to_physical for physical_id in row[2 : 2 + row[0]]) - assert all(physical_id == -1 for row in logical_to_physical for physical_id in row[2 + row[0] :]) + assert [row[1] for row in logical_to_physical] == [2, 2, 2, 2] + assert [row[2] for row in logical_to_physical] == [1, 1, 1, 1] + assert [row[3] for row in logical_to_physical] == [0, 1, 2, 3] + assert all(physical_id >= 0 for row in logical_to_physical for physical_id in row[3 : 3 + row[0]]) + assert all(physical_id == -1 for row in logical_to_physical for physical_id in row[3 + row[0] :]) def test_logical_to_physical_map_requires_expert_count_divisible_by_rank_count(): rank_to_logic_expert_ids = [[0, 1, 0], [2, 3, 1]] with pytest.raises(AssertionError): - build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=5, current_rank=0) + build_logical_to_physical_map( + rank_to_logic_expert_ids, + num_logical_experts=5, + current_rank=0, + node_world_size=2, + ) def test_logical_to_physical_map_supports_all_redundant_slots_for_one_expert(): @@ -585,42 +596,89 @@ def test_logical_to_physical_map_supports_all_redundant_slots_for_one_expert(): [[0, 1, 0, 0], [2, 3, 0, 0]], num_logical_experts=4, current_rank=0, + node_world_size=2, ) # 1 个主副本加上 2 个 rank 的全部 4 个冗余槽。 - assert logical_to_physical[0][0] == 5 - assert len(logical_to_physical[0][2:]) == 5 - assert len(set(logical_to_physical[0][2:])) == 5 + assert logical_to_physical[0][:3] == [5, 5, 3] + assert len(logical_to_physical[0][3:]) == 8 + assert len(set(logical_to_physical[0][3:8])) == 5 + assert logical_to_physical[0][8:] == [-1, -1, -1] def test_logical_to_physical_map_prefers_current_rank_replica(): redundant = [[4], [5], [0], [1]] rank_to_logic_expert_ids = _rank_to_logic_expert_ids(redundant, 8) - rank0_map = build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=8, current_rank=0) - rank1_map = build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=8, current_rank=1) + rank0_map = build_logical_to_physical_map( + rank_to_logic_expert_ids, + num_logical_experts=8, + current_rank=0, + node_world_size=2, + ) + rank1_map = build_logical_to_physical_map( + rank_to_logic_expert_ids, + num_logical_experts=8, + current_rank=1, + node_world_size=2, + ) fallback_redundant = [[4], [5], [0], [1], [2], [3]] - rank4_map = build_logical_to_physical_map(_rank_to_logic_expert_ids(fallback_redundant, 12), 12, current_rank=4) + rank4_map = build_logical_to_physical_map( + _rank_to_logic_expert_ids(fallback_redundant, 12), + 12, + current_rank=4, + node_world_size=2, + ) - assert rank0_map[0][0] == rank1_map[0][0] == 2 - assert rank0_map[0][1] == 1 - assert rank1_map[0][1] == 0 - assert rank0_map[0][2] == 0 - assert set(rank0_map[0][2:4]) == {0, 8} - assert set(rank1_map[0][2:4]) == {0, 8} - assert rank4_map[0][0] == 2 - assert rank4_map[0][1] == 0 - assert set(rank4_map[0][2:4]) == {0, 8} + assert rank0_map[0][:3] == [2, 1, 1] + assert rank1_map[0][:3] == [2, 1, 0] + assert rank0_map[0][3] == 0 + assert set(rank0_map[0][3:5]) == {0, 8} + assert set(rank1_map[0][3:5]) == {0, 8} + assert rank4_map[0][:3] == [2, 0, 0] + assert set(rank4_map[0][3:5]) == {0, 8} -def test_nonlocal_rank_routes_across_all_replicas(): +def test_nonlocal_rank_without_same_node_replica_routes_across_all_replicas(): redundant = [[4], [5], [0], [1]] rank_to_logic_expert_ids = _rank_to_logic_expert_ids(redundant, 8) - maps = [build_logical_to_physical_map(rank_to_logic_expert_ids, 8, current_rank=rank) for rank in range(4)] - assert maps[0][0][1] == 1 - assert maps[1][0][1] == 0 - assert maps[2][0][1] == 1 - assert maps[3][0][1] == 0 - assert all(set(logical_map[0][2:4]) == {0, 8} for logical_map in maps) + maps = [ + build_logical_to_physical_map( + rank_to_logic_expert_ids, + 8, + current_rank=rank, + node_world_size=1, + ) + for rank in range(4) + ] + assert maps[0][0][:3] == [2, 1, 1] + assert maps[1][0][:3] == [2, 0, 0] + assert maps[2][0][:3] == [2, 1, 1] + assert maps[3][0][:3] == [2, 0, 0] + assert all(set(logical_map[0][3:5]) == {0, 8} for logical_map in maps) + + +def test_nonlocal_rank_prefers_same_node_replica(): + rank_to_logic_expert_ids = [ + [1, 2], + [0, 3], + [0, 4], + [5, 6], + [0, 7], + [1, 2], + [3, 4], + [5, 6], + ] + + rank0_map = build_logical_to_physical_map( + rank_to_logic_expert_ids, + num_logical_experts=8, + current_rank=0, + node_world_size=4, + ) + + # Expert 0 is on same-node ranks 1 and 2 (physical IDs 2 and 4), plus + # remote rank 4 (physical ID 8). Hash routing only uses the first two. + assert rank0_map[0][:6] == [3, 2, 0, 2, 4, 8] def test_current_rank_moves_local_replica_to_front_without_changing_copies(): @@ -628,11 +686,21 @@ def test_current_rank_moves_local_replica_to_front_without_changing_copies(): # 自己的本地副本,因此路由槽的起点不同,但候选集合和数量保持一致。 redundant = [[1], [0], [3], [2]] rank_to_logic_expert_ids = _rank_to_logic_expert_ids(redundant, 4) - rank0_map = build_logical_to_physical_map(rank_to_logic_expert_ids, 4, current_rank=0) - rank1_map = build_logical_to_physical_map(rank_to_logic_expert_ids, 4, current_rank=1) + rank0_map = build_logical_to_physical_map( + rank_to_logic_expert_ids, + 4, + current_rank=0, + node_world_size=1, + ) + rank1_map = build_logical_to_physical_map( + rank_to_logic_expert_ids, + 4, + current_rank=1, + node_world_size=1, + ) - assert rank0_map[0] == [2, 1, 0, 3, -1, -1, -1] - assert rank1_map[0] == [2, 1, 3, 0, -1, -1, -1] + assert rank0_map[0] == [2, 1, 1, 0, 3, -1, -1, -1, -1, -1, -1] + assert rank1_map[0] == [2, 1, 1, 3, 0, -1, -1, -1, -1, -1, -1] def test_current_rank_stably_moves_all_local_physical_ids_to_front(): @@ -641,9 +709,14 @@ def test_current_rank_stably_moves_all_local_physical_ids_to_front(): # 仍保持原来的 [0, 1, 5] 顺序。 rank_to_logic_expert_ids = [[0, 0], [1, 0], [2, 0]] - rank1_map = build_logical_to_physical_map(rank_to_logic_expert_ids, num_logical_experts=3, current_rank=1) + rank1_map = build_logical_to_physical_map( + rank_to_logic_expert_ids, + num_logical_experts=3, + current_rank=1, + node_world_size=1, + ) - assert rank1_map[0] == [4, 1, 3, 0, 1, 5] + assert rank1_map[0] == [4, 1, 1, 3, 0, 1, 5, -1, -1] def test_transfer_plan_respects_explicit_target_slots(): @@ -838,6 +911,7 @@ def test_eplb_route_counter_has_one_entry_per_logical_expert(monkeypatch): monkeypatch.setattr(deepgemm_module, "get_env_start_args", lambda: args) monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(deepgemm_module, "get_node_world_size", lambda: 2) monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) original_zeros = torch.zeros @@ -964,6 +1038,7 @@ def repair(**kwargs): assert result[2].tolist() == [[128, 143]] assert repairs[0]["logical_topk_ids"] is logical_ids + assert repairs[0]["mode"] == "current_gpu_first" assert calls[0]["num_experts"] == 144 @@ -1035,6 +1110,7 @@ def repair(**kwargs): assert qinput == "qinput" assert calls[0]["logical_topk_ids"] is logical_ids assert not calls[0]["update_logical_expert_counter"] + assert calls[0]["mode"] == "current_gpu_first" def test_eplb_prefill_dispatch_consumes_physical_ids_and_event(monkeypatch): @@ -1101,6 +1177,7 @@ def repair(**kwargs): assert len(repair_calls) == 1 assert repair_calls[0]["logical_topk_ids"] is logical_ids assert repair_calls[0]["update_logical_expert_counter"] + assert repair_calls[0]["mode"] == "current_gpu_first" assert calls[0]["topk_idx"] is physical_ids assert calls[0]["topk_idx"].dtype is torch.long assert calls[0]["previous_event"] is caller_event @@ -1152,6 +1229,7 @@ def test_deepgemm_constructor_owns_eplb_runtime(monkeypatch): ) monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(deepgemm_module, "get_node_world_size", lambda: 2) monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) original_zeros = torch.zeros @@ -1183,6 +1261,7 @@ def test_deepgemm_constructor_loads_saved_layout_before_weight_initialization(mo ) monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(deepgemm_module, "get_node_world_size", lambda: 2) monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) monkeypatch.setattr( deepgemm_module, @@ -1210,7 +1289,7 @@ def test_deepgemm_constructor_loads_saved_layout_before_weight_initialization(mo impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace(), layer_index=7) assert impl.local_logics_expert_ids_list == saved_placement[0] - expected_map = build_logical_to_physical_map(saved_placement, 4, current_rank=0) + expected_map = build_logical_to_physical_map(saved_placement, 4, current_rank=0, node_world_size=2) assert impl.logical_to_physical_map.tolist() == expected_map @@ -1226,6 +1305,7 @@ def test_deepgemm_keeps_route_recording_when_rebalance_count_is_zero(monkeypatch ) monkeypatch.setattr(deepgemm_module, "get_global_world_size", lambda: 2) monkeypatch.setattr(deepgemm_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(deepgemm_module, "get_node_world_size", lambda: 2) monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) monkeypatch.setattr( deepgemm_module.torch, @@ -1257,6 +1337,7 @@ def repair(**kwargs): assert selected is physical_ids assert calls[0]["logical_topk_ids"] is logical_ids assert calls[0]["update_logical_expert_counter"] + assert calls[0]["mode"] == "current_gpu_first" def test_decode_masked_group_gemm_uses_all_physical_rows_when_eplb_is_enabled( @@ -1441,6 +1522,7 @@ def record_copy(tensor, source, non_blocking=False): target_placement[0], 6, current_rank=0, + node_world_size=2, ), dtype=torch.int32, ) @@ -1448,6 +1530,7 @@ def record_copy(tensor, source, non_blocking=False): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.global_rank = 0 manager.world_size = 2 + manager.node_world_size = 2 manager.num_logical_experts = 6 manager.target_placement = target_placement manager.current_placement = [[[0, 1, 2, 3, 2], [3, 4, 5, 0, 2]]] @@ -1680,6 +1763,7 @@ def is_finished(self): manager.active_transfer_batch = [transfer_info] manager.control_group = object() manager.world_size = 1 + manager.node_world_size = 1 manager.pending_transfer_batches = [] manager.num_logical_experts = 1 manager.global_rank = 0 @@ -2254,6 +2338,7 @@ def test_manager_initializes_without_transfer_task(monkeypatch): monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(manager_module, "get_node_world_size", lambda: 2) monkeypatch.setattr(manager_module, "get_eplb_step_interval", lambda: 20) clear_calls = [] monkeypatch.setattr( @@ -2311,7 +2396,8 @@ def all_gather_object(output, local_expert_ids_by_layer, group): @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") @pytest.mark.parametrize("update_logical_expert_counter", [False, True]) @pytest.mark.parametrize("tokens", [1, 32]) -def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tokens): +@pytest.mark.parametrize("mode", ["current_gpu_first", "current_node_first", "global_first"]) +def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tokens, mode): from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( eplb_repair_topk_ids, ) @@ -2323,16 +2409,35 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok logical_experts = torch.arange(experts, dtype=torch.int32, device="cuda") replica_counts = torch.where( logical_experts % 3 == 0, - torch.full_like(logical_experts, 2), + torch.full_like(logical_experts, 3), torch.ones_like(logical_experts), ) - has_local_replica = (logical_experts % 5 == 0).to(torch.int32) + num_current_gpu_replicas = torch.where( + logical_experts % 5 == 0, + torch.ones_like(logical_experts), + torch.zeros_like(logical_experts), + ) + num_node_replicas = torch.where( + num_current_gpu_replicas > 0, + torch.where( + replica_counts >= 2, + torch.full_like(logical_experts, 2), + num_current_gpu_replicas, + ), + torch.where( + (logical_experts % 7 == 0) & (replica_counts == 3), + torch.full_like(logical_experts, 2), + torch.zeros_like(logical_experts), + ), + ) logical_to_physical = torch.stack( ( replica_counts, - has_local_replica, + num_node_replicas, + num_current_gpu_replicas, logical_experts, - torch.where(replica_counts == 2, logical_experts + experts, logical_experts), + torch.where(replica_counts >= 2, logical_experts + experts, logical_experts), + torch.where(replica_counts == 3, logical_experts + 2 * experts, logical_experts), ), dim=1, ) @@ -2341,12 +2446,21 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok logical_ids_long = logical_ids.to(torch.long) token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) + if mode == "current_gpu_first": + num_preferred_replicas = torch.where( + num_current_gpu_replicas > 0, + num_current_gpu_replicas, + replica_counts, + ) + elif mode == "current_node_first": + num_preferred_replicas = torch.where(num_node_replicas > 0, num_node_replicas, replica_counts) + else: + num_preferred_replicas = replica_counts replica_indices = ( (((token_indices * 2654435769) & 0xFFFFFFFF) + ((logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF)) & 0xFFFFFFFF - ) % replica_counts[logical_ids_long].to(torch.int64) - replica_indices = torch.where(has_local_replica[logical_ids_long] != 0, 0, replica_indices) - expected_ids = logical_to_physical[logical_ids_long, replica_indices + 2] + ) % num_preferred_replicas[logical_ids_long].to(torch.int64) + expected_ids = logical_to_physical[logical_ids_long, replica_indices + 3] if update_logical_expert_counter: expected_counter.scatter_add_( 0, @@ -2359,6 +2473,7 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok logical_to_physical_map=logical_to_physical, logical_expert_counter=counter, update_logical_expert_counter=update_logical_expert_counter, + mode=mode, ) torch.cuda.synchronize() @@ -2378,6 +2493,7 @@ def test_eplb_repair_topk_ids_empty_input_skips_kernel(): counter = torch.zeros((experts,), dtype=torch.int64, device="cuda") logical_to_physical = torch.stack( ( + torch.ones((experts,), dtype=torch.int32, device="cuda"), torch.ones((experts,), dtype=torch.int32, device="cuda"), torch.ones((experts,), dtype=torch.int32, device="cuda"), torch.arange(experts, dtype=torch.int32, device="cuda"), @@ -2389,8 +2505,29 @@ def test_eplb_repair_topk_ids_empty_input_skips_kernel(): logical_to_physical_map=logical_to_physical, logical_expert_counter=counter, update_logical_expert_counter=True, + mode="current_gpu_first", ) assert physical_ids.shape == (0, 4) assert physical_ids.dtype is torch.int32 assert torch.equal(counter, torch.zeros_like(counter)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +def test_eplb_repair_topk_ids_rejects_unknown_dispatch_mode(): + from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( + eplb_repair_topk_ids, + ) + + logical_ids = torch.empty((0, 1), dtype=torch.int32, device="cuda") + logical_to_physical = torch.tensor([[1, 1, 1, 0]], dtype=torch.int32, device="cuda") + counter = torch.zeros((1,), dtype=torch.int64, device="cuda") + + with pytest.raises(AssertionError, match="unsupported EPLB dispatch mode"): + eplb_repair_topk_ids( + logical_topk_ids=logical_ids, + logical_to_physical_map=logical_to_physical, + logical_expert_counter=counter, + update_logical_expert_counter=False, + mode="unknown", + ) From afbe8df54a905ddb77adf92cef0cdef9afeef311 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 09:48:53 +0000 Subject: [PATCH 60/72] fix(eplb): improve replica hash distribution --- .../triton_kernel/fused_moe/eplb_topk_ids.py | 14 ++- unit_tests/common/fused_moe/test_eplb.py | 97 ++++++++++++++++++- 2 files changed, 104 insertions(+), 7 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py index c26d7eadc3..029f2fb48e 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py @@ -5,9 +5,17 @@ @triton.jit def _replica_index(token_index, logical_expert_id, num_valid_replicas): - token_hash = token_index.to(tl.uint32) * 2654435769 - expert_hash = logical_expert_id.to(tl.uint32) * 2246822519 - return (token_hash + expert_hash) % num_valid_replicas.to(tl.uint32) + # 先用 logical expert ID 给 token index 加盐,再用 32-bit avalanche + # finalizer 打散规律性 token 间隔,避免低位周期与副本数产生相关性。 + value = token_index.to(tl.uint32) + value ^= (logical_expert_id.to(tl.uint32) + 1) * 0x9E3779B9 + value ^= value >> 16 + value *= 0x7FEB352D + value ^= value >> 15 + value *= 0x846CA68B + value ^= value >> 16 + value = value.to(tl.uint32) + return value % num_valid_replicas.to(tl.uint32) @triton.jit diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index e9c55637d3..98b57caf2c 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -2456,10 +2456,14 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok num_preferred_replicas = torch.where(num_node_replicas > 0, num_node_replicas, replica_counts) else: num_preferred_replicas = replica_counts - replica_indices = ( - (((token_indices * 2654435769) & 0xFFFFFFFF) + ((logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF)) - & 0xFFFFFFFF - ) % num_preferred_replicas[logical_ids_long].to(torch.int64) + hash_values = token_indices ^ ((logical_ids.to(torch.int64) + 1) * 0x9E3779B9) + hash_values &= 0xFFFFFFFF + hash_values ^= hash_values >> 16 + hash_values = (hash_values * 0x7FEB352D) & 0xFFFFFFFF + hash_values ^= hash_values >> 15 + hash_values = (hash_values * 0x846CA68B) & 0xFFFFFFFF + hash_values ^= hash_values >> 16 + replica_indices = hash_values % num_preferred_replicas[logical_ids_long].to(torch.int64) expected_ids = logical_to_physical[logical_ids_long, replica_indices + 3] if update_logical_expert_counter: expected_counter.scatter_add_( @@ -2482,6 +2486,91 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok assert torch.equal(counter, expected_counter) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +def test_eplb_repair_topk_ids_spreads_strided_expert_tokens(): + from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( + eplb_repair_topk_ids, + ) + + num_tokens = 4096 + token_indices = torch.arange(num_tokens, dtype=torch.int32, device="cuda") + logical_ids = torch.where(token_indices % 4 == 0, 0, 1).view(-1, 1) + logical_to_physical = torch.tensor( + [ + [4, 4, 4, 0, 2, 3, 4], + [1, 1, 1, 1, -1, -1, -1], + ], + dtype=torch.int32, + device="cuda", + ) + counter = torch.zeros((2,), dtype=torch.int64, device="cuda") + + physical_ids = eplb_repair_topk_ids( + logical_topk_ids=logical_ids, + logical_to_physical_map=logical_to_physical, + logical_expert_counter=counter, + update_logical_expert_counter=False, + mode="global_first", + ) + torch.cuda.synchronize() + + strided_token_outputs = physical_ids[token_indices % 4 == 0, 0] + replica_counts = torch.stack([(strided_token_outputs == physical_id).sum() for physical_id in (0, 2, 3, 4)]) + assert torch.all(replica_counts > strided_token_outputs.numel() * 0.2) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +@pytest.mark.parametrize("num_replicas", [2, 3, 4, 5, 101, 127, 128, 251]) +def test_eplb_replica_hash_is_uniform_across_tokens_and_experts(num_replicas): + from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( + eplb_repair_topk_ids, + ) + + num_tokens = 4097 + num_experts = 256 + expert_ids = torch.arange(num_experts, dtype=torch.int32, device="cuda") + logical_ids = expert_ids.expand(num_tokens, -1).contiguous() + replica_counts = torch.full((num_experts, 3), num_replicas, dtype=torch.int32, device="cuda") + physical_ids = expert_ids.unsqueeze(1) + num_experts * torch.arange( + num_replicas, + dtype=torch.int32, + device="cuda", + ) + logical_to_physical = torch.cat((replica_counts, physical_ids), dim=1) + counter = torch.zeros((num_experts,), dtype=torch.int64, device="cuda") + + routed_physical_ids = eplb_repair_topk_ids( + logical_topk_ids=logical_ids, + logical_to_physical_map=logical_to_physical, + logical_expert_counter=counter, + update_logical_expert_counter=False, + mode="global_first", + ) + torch.cuda.synchronize() + + routed_replica_indices = routed_physical_ids // num_experts + observed_counts = torch.stack( + [(routed_replica_indices == replica_index).sum(dim=0) for replica_index in range(num_replicas)], + dim=1, + ) + expected_count = num_tokens / num_replicas + deviations = observed_counts - expected_count + if num_replicas <= 5: + max_relative_deviation = (deviations.abs() / expected_count).max().item() + assert max_relative_deviation < 0.12 + else: + # 副本很多时单槽期望样本较少,使用每个 expert 的归一化卡方值 + # 检查整体形状,并额外检查跨 expert 汇总后的单槽最大偏差。 + normalized_chi_square = (deviations.square() / expected_count).sum(dim=1) / (num_replicas - 1) + assert normalized_chi_square.max().item() < 1.75 + + aggregate_expected_count = num_tokens * num_experts / num_replicas + aggregate_max_relative_deviation = ( + (observed_counts.sum(dim=0) - aggregate_expected_count).abs() / aggregate_expected_count + ).max() + assert aggregate_max_relative_deviation.item() < 0.06 + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") def test_eplb_repair_topk_ids_empty_input_skips_kernel(): from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( From d20f3f06a7d217d4039baf68267ac8c84b8d4d34 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 09:55:28 +0000 Subject: [PATCH 61/72] refactor(eplb): simplify MoE weight discovery --- .../mode_backend/eplb/runtime_manager.py | 12 ++++++------ unit_tests/common/fused_moe/test_eplb.py | 19 +++++++++---------- 2 files changed, 15 insertions(+), 16 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index ac0f531b5b..715e6755f4 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -1,6 +1,6 @@ from enum import Enum import time -from typing import Dict, List, Optional +from typing import List, Optional import torch import torch.distributed as dist @@ -479,12 +479,12 @@ def _persist_current_placement(self) -> None: def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: - weights_by_id: Dict[int, FusedMoeWeight] = {} + weights: List[FusedMoeWeight] = [] for layer in model.trans_layers_weight: - for value in getattr(layer, "__dict__", {}).values(): - if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: - weights_by_id[id(value)] = value - return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) + weight = getattr(layer, "experts", None) + if isinstance(weight, FusedMoeWeight) and weight.enable_ep_moe: + weights.append(weight) + return weights def _expert_load_imbalance_ratio(expert_load: torch.Tensor) -> float: diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 98b57caf2c..f45c83d6bc 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -215,27 +215,26 @@ def test_factory_selects_all_paths_without_ep_constructor_state(monkeypatch): ) -def test_find_fused_moe_weights_discovers_direct_layer_attributes(monkeypatch): +def test_find_fused_moe_weights_uses_layer_experts_in_model_order(monkeypatch): class FakeFusedMoeWeight: def __init__(self, layer_num, enable_ep_moe=True): self.layer_num_ = layer_num self.enable_ep_moe = enable_ep_moe monkeypatch.setattr(manager_module, "FusedMoeWeight", FakeFusedMoeWeight) - first = FakeFusedMoeWeight(3) - alternate = FakeFusedMoeWeight(1) - aliased = FakeFusedMoeWeight(2) - disabled = FakeFusedMoeWeight(0, enable_ep_moe=False) + first = FakeFusedMoeWeight(1) + second = FakeFusedMoeWeight(3) + disabled = FakeFusedMoeWeight(2, enable_ep_moe=False) model = SimpleNamespace( trans_layers_weight=[ - SimpleNamespace(moe_weight=first), - SimpleNamespace(alternate_moe_weight=alternate), - SimpleNamespace(moe_weight=aliased, alternate_moe_weight=aliased), - SimpleNamespace(moe_weight=disabled), + SimpleNamespace(experts=first), + SimpleNamespace(), + SimpleNamespace(experts=disabled), + SimpleNamespace(experts=second), ] ) - assert manager_module._find_fused_moe_weights(model) == [alternate, aliased, first] + assert manager_module._find_fused_moe_weights(model) == [first, second] def test_eplb_redundant_experts_default_to_disabled(): From 4e774e7808a1b82405a0eee2fca3a958384d3020 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 10:13:02 +0000 Subject: [PATCH 62/72] feat(eplb): support configurable placement planners --- docs/CN/source/tutorial/api_server_args.rst | 13 ++++ docs/EN/source/tutorial/api_server_args.rst | 14 ++++ lightllm/server/api_cli.py | 8 ++ lightllm/server/core/objs/start_args_type.py | 1 + .../model_infer/mode_backend/base_backend.py | 1 + .../mode_backend/eplb/placement/__init__.py | 2 + .../mode_backend/eplb/placement/factory.py | 25 ++++++ .../mode_backend/eplb/runtime_manager.py | 76 ++++++++++++++----- unit_tests/common/fused_moe/test_eplb.py | 26 +++++++ 9 files changed, 147 insertions(+), 19 deletions(-) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/placement/factory.py diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index ad21c4a53d..2c291f94d0 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -724,6 +724,18 @@ PD 分离模式参数 EPLB 当前不能与 ``--enable_prefill_cudagraph`` 同时使用,也不支持 SM100 GPU。 同一部署中的所有 rank 和节点必须使用相同的配置值。 +.. option:: --eplb_plan_mode + + EPLB 动态重排使用的专家布局规划算法,默认值为 ``greedy``。当前支持: + + * ``greedy``:根据各层逻辑专家的全局路由负载生成近似均衡的完整布局, + 并尽量复用当前 rank 和物理槽位以减少专家迁移。 + + 此参数只选择布局规划算法,不改变 token 到已有专家副本的运行时分发策略。 + 同一个 EP 通信组内的所有 rank 必须使用相同的值。PD 分离部署中的 prefill + 和 decode 进程拥有各自独立的 EPLB manager,因此可以分别设置适合各自流量 + 特征的规划算法;非 PD 部署则使用一个算法处理该进程采集到的全部路由负载。 + .. option:: --eplb_rebalance_count 动态 EPLB 最多成功执行的重排次数,默认值为 ``1``。只有新布局实际发生 @@ -749,6 +761,7 @@ PD 分离模式参数 --model_dir /path/to/model \ --enable_ep_moe \ --eplb_num_redundant_experts_per_rank 2 \ + --eplb_plan_mode greedy \ --eplb_rebalance_count 1 \ --eplb_config_path /path/to/eplb-placement.json diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 83a671f175..6f3a07fe1a 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -742,6 +742,19 @@ Expert Parallelism and EPLB Parameters EPLB currently cannot be combined with ``--enable_prefill_cudagraph`` and is not supported on SM100 GPUs. Use the same value on every rank and node in one deployment. +.. option:: --eplb_plan_mode + + Expert placement planning algorithm used for dynamic EPLB rebalances. The default is ``greedy``. The + currently supported value is: + + * ``greedy``: builds an approximately balanced full placement from the global logical-expert load of each + layer and attempts to reuse the current ranks and physical slots to reduce expert migration. + + This option selects the placement planner; it does not change how tokens are dispatched among replicas in + an existing placement. Every rank in one EP communication group must use the same value. Prefill and decode + processes in a PD-disaggregated deployment have independent EPLB managers and may select planners suited to + their respective traffic. A non-PD process uses one planner for all routing load collected by that process. + .. option:: --eplb_rebalance_count Maximum number of successfully completed dynamic EPLB rebalances. The default is ``1``. A count is consumed @@ -769,6 +782,7 @@ Expert Parallelism and EPLB Parameters --model_dir /path/to/model \ --enable_ep_moe \ --eplb_num_redundant_experts_per_rank 2 \ + --eplb_plan_mode greedy \ --eplb_rebalance_count 1 \ --eplb_config_path /path/to/eplb-placement.json diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 786558be30..646cb4d5b3 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -773,6 +773,14 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: help="""Number of redundant physical experts per EP rank for each MoE layer. Set to 0 to disable EPLB.""", ) + parser.add_argument( + "--eplb_plan_mode", + type=str, + choices=["greedy"], + default="greedy", + help="""EPLB placement planning algorithm used by this inference process. + Prefill and decode processes may select their planner independently in PD deployments.""", + ) parser.add_argument( "--eplb_rebalance_count", type=int, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 8f7efc08c8..a3ff8061df 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -187,6 +187,7 @@ class StartArgs: ) enable_ep_moe: bool = field(default=False) eplb_num_redundant_experts_per_rank: int = field(default=0) + eplb_plan_mode: str = field(default="greedy", metadata={"choices": ["greedy"]}) eplb_rebalance_count: int = field(default=1) eplb_config_path: Optional[str] = field(default=None) enable_fused_shared_experts: bool = field(default=False) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index ffcf2b0e3a..48de845967 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -264,6 +264,7 @@ def init_model(self, kvargs): self.model, max_rebalance_count=self.args.eplb_rebalance_count, config_path=self.args.eplb_config_path, + plan_mode=self.args.eplb_plan_mode, ) # 启动infer_loop_thread, 启动两个线程进行推理,对于具备双batch推理折叠得场景 diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py index 8b68634d0e..25df90a986 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py @@ -5,6 +5,7 @@ from .initial import build_initial_local_expert_ids from .routing import build_logical_to_physical_map from .greedy import GreedyEPLBPlanner +from .factory import create_eplb_planner from .config import load_layer_placement, save_placement_config __all__ = [ @@ -16,6 +17,7 @@ "LogicalToPhysicalMap", "build_initial_local_expert_ids", "build_logical_to_physical_map", + "create_eplb_planner", "load_layer_placement", "save_placement_config", ] diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/factory.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/factory.py new file mode 100644 index 0000000000..97c9e24dcd --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/factory.py @@ -0,0 +1,25 @@ +"""EPLB placement planner selection.""" + +from typing import Callable, Dict + +from .greedy import GreedyEPLBPlanner +from .planner import EPLBPlanner + + +def create_eplb_planner( + plan_mode: str, + num_ranks: int, + num_redundant_experts_per_rank: int, + expert_alignment: int, +) -> EPLBPlanner: + """根据启动参数为当前推理进程创建布局规划器。""" + planner_builders: Dict[str, Callable[[], EPLBPlanner]] = { + "greedy": lambda: GreedyEPLBPlanner( + num_ranks, + num_redundant_experts_per_rank, + expert_alignment=expert_alignment, + ), + } + if plan_mode not in planner_builders: + raise ValueError(f"unsupported EPLB plan mode {plan_mode!r}; expected one of {tuple(planner_builders)}") + return planner_builders[plan_mode]() diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 715e6755f4..fb964d2409 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -28,8 +28,8 @@ from .placement import ( EPLBPlanner, ExpertPlacement, - GreedyEPLBPlanner, build_logical_to_physical_map, + create_eplb_planner, save_placement_config, ) from .placement_plan_task import EPLBPlanTask @@ -55,18 +55,53 @@ class EPLBManagerState(Enum): class EPLBManager: """由 :meth:`step` 驱动的 EPLB 状态机。 - 状态循环如下: - - ``COLLECTING -> EVALUATING -> PLAN_PLACEMENT -> WAIT_PLAN_PLACEMENT_FINISHED`` - ``-> PLAN_TRANSFER -> WAIT_PLAN_TRANSFER_FINISHED -> TRANSFERRING`` - - 当累计的平均专家 token 数不足时,``EVALUATING`` 会回到 - ``COLLECTING``;当规划器认为无需调整布局时, - ``WAIT_PLAN_PLACEMENT_FINISHED`` 会 - 回到 ``COLLECTING``。达到重排次数上限后,manager 仍会周期性采集并 - 上报本地负载指标,但不再进入全局负载汇总和布局规划。每次调用 - :meth:`step` 最多推进一个状态,布局规划、传输规划和权重传输都在后台 - 执行,主推理线程负责评估、轮询和提交结果。 + 每次调用 :meth:`step` 最多处理一个状态。主路径及各状态的职责如下:: + + [COLLECTING] + 推理 kernel 持续累计各层 logical expert 的 route counter; + manager 只记录采样 step,等待下一个评估周期。 + | + | 评估周期到达 + v + [EVALUATING] + 将本地 counter 快照到 CPU 并上报负载指标;汇总各 rank 的 + token 总数,判断样本量和剩余重排次数。 + | + | 样本充足且仍允许重排 + v + [PLAN_PLACEMENT] + 汇集完整的全局专家负载;rank 0 启动后台布局规划任务。 + | + v + [WAIT_PLAN_PLACEMENT_FINISHED] + 轮询 rank 0 的规划任务,并向所有 rank 广播目标布局。 + | + | 目标布局发生变化 + v + [PLAN_TRANSFER] + 每个 rank 根据相同的当前/目标布局启动后台传输规划任务。 + | + v + [WAIT_PLAN_TRANSFER_FINISHED] + 等待所有 rank 生成一致的、按依赖关系分批的传输任务。 + | + v + [TRANSFERRING] + 分批启动并轮询后台权重传输;整批完成后,主推理线程在安全 + 边界统一提交权重和路由 metadata。全部批次完成后发布新布局、 + 清空 counter 并回到 COLLECTING。 + + 以下分支会提前回到 ``COLLECTING``:: + + EVALUATING + |-- 平均 token 数不足 --------> 保留 counter,继续累计样本 + `-- 已达到重排次数上限 ------> 清空 counter,仅周期性上报指标 + + WAIT_PLAN_PLACEMENT_FINISHED + `-- 目标布局与当前布局相同 ---> 保留 counter,等待下次评估 + + 布局规划、传输规划和权重传输在后台执行;主推理线程只负责创建任务、 + 轮询状态,以及在安全边界提交已经完成的结果。 """ def __init__( @@ -74,6 +109,7 @@ def __init__( model: TpPartBaseModel, max_rebalance_count: int = 1, config_path: Optional[str] = None, + plan_mode: str = "greedy", ) -> None: # SM100 FP4 Mega-MoE 会将在线专家权重转换为独立的 kernel 布局,并使用源 tensor 的 data_ptr # 作为 key 缓存这些转换后的副本。EPLB 通过原地 copy_ 替换专家行,只改变权重内容而不会改变 @@ -98,6 +134,13 @@ def __init__( first_impl = self._eplb_impls[0] self.num_logical_experts: int = first_impl.n_routed_experts self.num_redundant_experts_per_rank: int = first_impl.num_redundant_experts_per_rank + self.plan_mode: str = plan_mode + self.planner: EPLBPlanner = create_eplb_planner( + self.plan_mode, + self.world_size, + self.num_redundant_experts_per_rank, + expert_alignment=EPLB_EXPERT_ALIGNMENT, + ) # 评估调度:steps 只在 COLLECTING 状态递增。route counter 从当前 # 布局生效时开始累计,让低流量服务可以跨多个评估周期收集足够样本。 @@ -128,11 +171,6 @@ def __init__( [expert_ids_by_rank_and_layer[rank][layer_index] for rank in range(self.world_size)] for layer_index in range(len(weights)) ] - self.planner: EPLBPlanner = GreedyEPLBPlanner( - self.world_size, - self.num_redundant_experts_per_rank, - expert_alignment=EPLB_EXPERT_ALIGNMENT, - ) self.state = EPLBManagerState.COLLECTING self.next_evaluation_step = self.step_interval @@ -144,7 +182,7 @@ def __init__( f"eplb enabled layers={len(weights)} num_logical_experts={self.num_logical_experts} " f"num_redundant_experts_per_rank={self.num_redundant_experts_per_rank} " f"step_interval={self.step_interval} max_rebalance_count={self.max_rebalance_count} " - f"planner={type(self.planner).__name__}" + f"plan_mode={self.plan_mode} planner={type(self.planner).__name__}" ) def step(self) -> None: diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index f45c83d6bc..cd85b8f1e3 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -11,6 +11,7 @@ GreedyEPLBPlanner, build_initial_local_expert_ids, build_logical_to_physical_map, + create_eplb_planner, ) from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs @@ -243,6 +244,9 @@ def test_eplb_redundant_experts_default_to_disabled(): assert parser.parse_args([]).eplb_num_redundant_experts_per_rank == 0 assert parser.parse_args(["--eplb_num_redundant_experts_per_rank", "3"]).eplb_num_redundant_experts_per_rank == 3 assert StartArgs().eplb_num_redundant_experts_per_rank == 0 + assert parser.parse_args([]).eplb_plan_mode == "greedy" + assert parser.parse_args(["--eplb_plan_mode", "greedy"]).eplb_plan_mode == "greedy" + assert StartArgs().eplb_plan_mode == "greedy" assert parser.parse_args([]).eplb_rebalance_count == 1 assert parser.parse_args(["--eplb_rebalance_count", "-1"]).eplb_rebalance_count == -1 assert parser.parse_args(["--eplb_rebalance_count", "0"]).eplb_rebalance_count == 0 @@ -286,6 +290,25 @@ def test_eplb_planner_defines_an_abstract_planning_interface(): assert isinstance(GreedyEPLBPlanner(2, 1), EPLBPlanner) +def test_create_eplb_planner_selects_requested_algorithm(): + planner = create_eplb_planner( + "greedy", + 2, + 1, + expert_alignment=1, + ) + + assert isinstance(planner, GreedyEPLBPlanner) + + with pytest.raises(ValueError, match="unsupported EPLB plan mode"): + create_eplb_planner( + "unknown", + 2, + 1, + expert_alignment=1, + ) + + def test_eplb_planner_builds_legal_concrete_slot_layout(): planner = GreedyEPLBPlanner( 4, @@ -2386,7 +2409,10 @@ def all_gather_object(output, local_expert_ids_by_layer, group): assert manager.next_evaluation_step == manager.step_interval assert manager.max_rebalance_count == 1 assert manager.completed_rebalance_count == 0 + assert manager.plan_mode == "greedy" assert clear_calls == [manager] + assert isinstance(manager.planner, GreedyEPLBPlanner) + assert "plan_mode=greedy" in logs[0] assert "planner=GreedyEPLBPlanner" in logs[0] assert weight.fuse_moe_impl.recording assert manager._eplb_impls[0] is weight.fuse_moe_impl From 2680dec2bed0be80822cd6a4cbd2bf134f987697 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 21 Sep 2026 10:20:51 +0000 Subject: [PATCH 63/72] docs(eplb): add Chinese implementation guide --- docs/CN/source/framework/eplb.md | 404 +++++++++++++++++++++++++++++++ docs/CN/source/index.rst | 1 + 2 files changed, 405 insertions(+) create mode 100644 docs/CN/source/framework/eplb.md diff --git a/docs/CN/source/framework/eplb.md b/docs/CN/source/framework/eplb.md new file mode 100644 index 0000000000..578150f095 --- /dev/null +++ b/docs/CN/source/framework/eplb.md @@ -0,0 +1,404 @@ +# EPLB 专家负载均衡实现 + +本文介绍 LightLLM 中 Expert Parallelism Load Balancer(EPLB)的完整实现,包括物理专家槽位、在线路由、负载采集、布局规划、权重迁移、运行时状态机、布局持久化,以及如何扩展新的规划算法。 + +EPLB 的目标是在不改变模型逻辑专家语义的前提下,利用额外的物理专家副本缓解热点专家造成的 EP rank 负载不均。它把“模型选择了哪个逻辑专家”和“本次由哪个物理副本执行”分成两个阶段,并允许服务运行期间重新安排物理副本。 + +## 1. 核心概念 + +设: + +- `E`:模型每层的逻辑专家数; +- `W`:EP world size; +- `R`:每个 rank 配置的冗余专家数; +- `E / W`:每个 rank 原本持有的专家数; +- `E / W + R`:每个 rank 实际分配的物理专家槽位数; +- `E + W * R`:一个 MoE 层在整个 EP world 中的物理槽位总数。 + +逻辑专家 ID 来自模型路由器,范围固定为 `[0, E)`。物理专家 ID 标识实际执行权重所在的槽位: + +```text +physical_expert_id = rank * num_physical_experts_per_rank + local_slot +``` + +一个逻辑专家可以拥有多个物理副本,但同一个 rank 上不会重复放置同一个逻辑专家。任意合法布局还必须满足: + +1. 每个 rank 的物理槽位数相同; +2. 所有逻辑专家至少有一个物理副本; +3. 所有逻辑专家 ID 都在 `[0, E)` 范围内; +4. 同一 rank 内的逻辑专家 ID 不重复。 + +## 2. 启用方式 + +EPLB 通过冗余专家数量开启: + +```bash +python -m lightllm.server.api_server \ + --model_dir /path/to/model \ + --enable_ep_moe \ + --eplb_num_redundant_experts_per_rank 2 \ + --eplb_plan_mode greedy \ + --eplb_rebalance_count 1 \ + --eplb_config_path /path/to/eplb-placement.json +``` + +主要参数如下: + +| 参数 | 默认值 | 作用 | +| --- | --- | --- | +| `--enable_ep_moe` | 关闭 | 启用专家并行;EPLB 的前置条件 | +| `--eplb_num_redundant_experts_per_rank` | `0` | 每个 rank 的额外物理专家槽位数;大于 0 时启用 EPLB | +| `--eplb_plan_mode` | `greedy` | 选择动态布局规划算法;当前支持 `greedy` | +| `--eplb_rebalance_count` | `1` | 最多完成的动态重排次数;`-1` 表示不限次数,`0` 表示不动态重排 | +| `--eplb_config_path` | `None` | 可选的布局加载与回写路径 | + +完整命令行说明见 {doc}`../tutorial/api_server_args`。 + +在 PD 分离部署中,prefill 和 decode 进程各自拥有独立的 EPLB manager,可以分别设置 `--eplb_plan_mode`。同一个 EP 通信组内的所有 rank 必须使用相同配置。非 PD 部署只有一个 manager,它根据该进程采集到的全部路由负载生成统一布局。 + +## 3. 总体架构 + +```text +模型路由器 + │ + │ logical top-k IDs + v +EPLB 路由 kernel + ├── 按 logical expert 累计 route_counter + ├── 查询 logical_to_physical_map + └── 为每个 token 选择 physical expert ID + │ + v + MoE 执行 kernel + +周期性控制面: + +route_counter + -> 全局负载汇总 + -> placement planner + -> target placement + -> transfer planner + -> 后台权重传输 + -> 安全边界提交权重和路由 metadata +``` + +主要实现位置: + +| 模块 | 职责 | +| --- | --- | +| `fused_moe/impl/deepgemm_impl.py` | 初始化 EPLB 运行态,在 MoE 执行前修复 logical top-k IDs | +| `triton_kernel/fused_moe/eplb_topk_ids.py` | 统计逻辑专家负载,并把 logical ID 映射为 physical ID | +| `eplb/placement/initial.py` | 构建确定性的初始专家布局 | +| `eplb/placement/routing.py` | 根据完整布局构建紧凑路由表 | +| `eplb/placement/planner.py` | 布局规划器抽象接口 | +| `eplb/placement/factory.py` | 根据 `eplb_plan_mode` 创建具体规划器 | +| `eplb/placement/greedy.py` | 默认的贪心布局算法 | +| `eplb/async_transfer_planner.py` | 在后台生成跨层传输批次 | +| `eplb/expert_transfer.py` | 规划槽位依赖并执行专家权重传输 | +| `eplb/runtime_manager.py` | 驱动状态机,协调采集、规划、传输和提交 | + +## 4. 初始化布局与权重加载 + +### 4.1 默认布局 + +启动时首先把逻辑专家连续划分到各 rank,然后从下一个 rank 的主专家区间开始循环选择冗余副本。例如 `E=8`、`W=4`、`R=2` 时: + +```text +rank 0: [0, 1, 2, 3] +rank 1: [2, 3, 4, 5] +rank 2: [4, 5, 6, 7] +rank 3: [6, 7, 0, 1] +``` + +每行前两个槽位来自原始连续划分,后两个槽位是启动时已经加载完成的冗余副本。运行期允许重新分配所有物理槽位,不再区分不可移动的“主槽位”和只能替换的“冗余槽位”。 + +### 4.2 从历史布局启动 + +指定 `--eplb_config_path` 后,每个 MoE 层会尝试加载历史布局。配置必须同时匹配: + +- 配置版本; +- 逻辑专家数; +- world size; +- 每个 rank 的冗余专家数; +- 模型层号; +- 每层布局形状、专家 ID 范围、rank 内唯一性和全专家覆盖关系。 + +任意校验失败都会记录 warning,并仅对受影响的层回退到默认布局。校验通过时,专家权重会直接按照历史布局加载,不需要服务启动后再执行一次恢复迁移。 + +### 4.3 本地运行态 + +每个 MoE 实现对象持有: + +- `local_logics_expert_ids_list`:本 rank 每个物理槽对应的逻辑专家; +- `logical_to_physical_map`:logical ID 到可用 physical IDs 的设备路由表; +- `route_counter`:长度为 `E` 的 `int64` GPU 计数器; +- `num_redundant_experts_per_rank`:本 rank 的额外槽位数。 + +目前启用 EP MoE 时使用 `FuseMoeDeepGEMM` 实现。EPLB manager 只收集启用了 EP 的 `layer.experts`,并保留模型中的层顺序。 + +## 5. 在线路由与负载采集 + +### 5.1 logical ID 与 physical ID 分离 + +MoE 路由器首先只在模型的逻辑专家空间中计算 top-k: + +```text +_select_experts + -> topk_weights + logical_topk_ids + -> capture callback + -> _prepare_expert_execution + -> EPLB logical-to-physical 映射 + -> _fused_experts +``` + +逻辑 ID tensor 不会被原地修改。监控和 capture callback 始终看到模型语义上的 logical expert;只有实际执行 MoE kernel 前才生成新的 physical ID tensor。 + +### 5.2 路由表布局 + +每个 logical expert 对应一行固定宽度 metadata: + +```text +[global_count, node_count, current_gpu_count, + physical_ids..., -1 padding...] +``` + +- `global_count`:整个 EP world 中的有效副本数; +- `node_count`:当前节点内的有效副本数,包含本卡; +- `current_gpu_count`:当前 GPU 上的有效副本数; +- `physical_ids`:按“本卡、本节点其他卡、其他节点”的顺序稳定排列; +- `padding`:未使用槽位填 `-1`,kernel 不会读取。 + +路由槽位上限等于整个 world 的物理槽位总数,因此布局变化不会改变 tensor 的 shape。 + +### 5.3 副本分发模式 + +路由算子要求调用方显式指定分发模式: + +| 模式 | 候选副本 | +| --- | --- | +| `current_gpu_first` | 本卡存在副本时只在本卡副本间选择,否则回退到全局副本 | +| `current_node_first` | 本节点存在副本时在节点内选择,否则回退到全局副本 | +| `global_first` | 直接在全局全部有效副本间选择 | + +当前 DeepGEMM EPLB 路径使用 `current_gpu_first`。未来如果要支持“本卡 -> 本节点 -> 全局”的三级回退,需要布局规划算法同时具备节点拓扑感知能力。 + +### 5.4 副本哈希 + +同一候选集合内使用 `(token_index, logical_expert_id)` 生成 32 位哈希,再对有效副本数取模。实现先用 logical expert ID 给 token index 加盐,然后执行 32 位 avalanche finalizer。 + +该变换由奇数乘法和可逆的异或移位组成,可以显著降低规律性 token 间隔与副本数之间的低位相关性。例如同一专家每隔 4 个 token 出现且有 4 个副本时,简单线性哈希可能退化到单一副本,avalanche mix 能将流量重新打散。 + +### 5.5 负载统计 + +路由 kernel 在 physical ID 映射前,按 logical expert 对 `route_counter` 执行原子累加。这样同一逻辑专家的多个物理副本不会拆散规划器观察到的负载信号。 + +计数器与 MoE forward 位于同一条 overlap stream。清零操作也提交到该 stream,从而自然排在此前 forward 之后、后续 forward 之前,无需额外的全设备同步。 + +## 6. EPLB 状态机 + +`EPLBManager.step()` 在安全的推理边界被调用,每次最多处理一个状态: + +```text +[COLLECTING] + 累计 logical route counter,等待评估周期 + | + v +[EVALUATING] + CPU 快照、指标上报、样本量与重排次数检查 + | + v +[PLAN_PLACEMENT] + 汇集全局负载,rank 0 启动后台布局规划 + | + v +[WAIT_PLAN_PLACEMENT_FINISHED] + 轮询规划结果,并广播目标布局 + | + v +[PLAN_TRANSFER] + 各 rank 根据相同布局启动后台传输规划 + | + v +[WAIT_PLAN_TRANSFER_FINISHED] + 等待所有 rank 生成一致的传输批次 + | + v +[TRANSFERRING] + 分批启动/轮询权重传输,在安全边界统一提交 + | + `-------------------------------> COLLECTING +``` + +提前返回 `COLLECTING` 的分支: + +```text +EVALUATING + |-- 平均 token 数不足 ----------> 保留 counter,继续累计 + `-- 达到重排次数上限 ----------> 清空 counter,只做周期性指标上报 + +WAIT_PLAN_PLACEMENT_FINISHED + `-- 目标布局与当前布局相同 -----> 保留 counter,等待下次评估 +``` + +默认每 20 个采样 step 评估一次,可以通过环境变量 `LIGHTLLM_EPLB_STEP_INTERVAL` 调整。该值必须大于 0。 + +只有当整个 world 的平均样本量达到每个“层 × 逻辑专家”256 个 token 时才开始规划。样本不足不会清空 counter,低流量服务可以跨多个评估周期累计数据。 + +## 7. 布局规划 + +### 7.1 规划器接口与选择 + +所有布局算法实现统一的 `EPLBPlanner.plan(logical_expert_load, current_placement)` 接口,返回: + +```text +[layer][rank][local physical slot] -> logical expert ID +``` + +`--eplb_plan_mode` 只负责选择布局算法。`create_eplb_planner` 将字符串模式转换成具体实例,使状态机不依赖某个算法类。当前唯一模式为 `greedy`。 + +规划只在 rank 0 的后台线程执行。完成后,目标布局通过控制通信组广播给所有 rank。相同输入必须产生确定结果,便于所有 rank 生成一致的传输计划。 + +### 7.2 Greedy 规划算法 + +默认算法按层独立规划,主要步骤如下: + +1. **选择全卡冗余专家**:选取负载最高的 `R` 个逻辑专家,在每个 rank 上各放置一份; +2. **确定额外副本数**:其余专家先各保留一个副本,再把剩余 `R` 个副本逐次分给当前 `load / replica_count` 最大的专家; +3. **平铺多副本专家**:使用循环 rank 游标,把同一专家的副本放到不同 rank; +4. **放置单副本专家**:按专家负载从高到低处理,每次放到当前估算负载最低且仍有空槽的 rank; +5. **复用当前布局**:先把候选 rank 行匹配到共同专家最多的当前 rank,再让共同专家尽量保留原物理槽位,以减少跨 rank 传输和 rank 内覆盖。 + +规划负载按 128 token 对齐,降低很小的计数波动对布局的影响。专家 ID 和 rank ID 用作稳定的平局规则,因此结果是确定性的。 + +## 8. 权重迁移与安全提交 + +### 8.1 传输计划 + +传输规划器逐层比较当前布局和目标布局,为每个变化的目标槽绑定一个确定的源槽。选择源槽时优先使用不会被覆盖的稳定副本;没有稳定副本时,循环使用当前已有副本,避免把读取集中在同一个 rank。 + +随后根据“目标槽是否仍是其他任务的源槽”建立覆盖依赖: + +- **安全任务**:目标槽不再承担待处理任务的源,可以先传输并提交; +- **依赖环**:所有目标槽同时也是源槽,必须先把整个环的权重读入 pinned memory,再统一覆盖; +- **rank 冲突拆批**:普通批次中每个 rank 最多参与一条任务,在限制 pinned memory 峰值的同时保留跨 rank 并行性。 + +每层单独生成批次,再按层顺序拼接。这样完成一层的提交后就能立即发布该层的新路由 metadata。 + +### 8.2 数据路径 + +远程专家的传输路径为: + +```text +源 GPU 权重行 + -> 源 rank pinned CPU row + -> Gloo point-to-point + -> 目标 rank pinned CPU row + -> 目标 GPU live 权重行 +``` + +如果源和目标属于同一个 rank,则只执行 GPU 到 pinned CPU 的本地暂存,不经过网络。一次专家传输会覆盖实际推理需要的全部张量,包括量化权重及其 scale、zero point 等配套状态。 + +控制面和权重传输分别使用独立的 Gloo process group,避免两类通信相互干扰。后台线程只负责把数据传入 pinned memory,不直接修改 live 权重。 + +### 8.3 提交边界 + +只有当所有 rank 都确认当前批次传输完成后,主推理线程才会在 overlap stream 上统一: + +1. 把目标 rank 的 pinned row 写入 live GPU 权重槽; +2. 更新 `current_placement` 和本地槽位的 logical expert ID; +3. 为发生变化的层重建 `logical_to_physical_map`; +4. 将新 metadata 异步复制到 GPU。 + +权重和路由 metadata 在同一条 stream 上更新,后续 forward 只能看到完整提交后的状态,不会观察到“新路由指向旧权重”或“旧路由指向新权重”的中间状态。 + +全部批次完成后,manager 发布目标布局、清空 route counter、增加完成次数,并回到 `COLLECTING`。 + +## 9. 布局持久化 + +成功完成重排后,rank 0 会把最新完整布局写回 `--eplb_config_path`。配置内容包括: + +```json +{ + "version": 1, + "num_logical_experts": 8, + "world_size": 4, + "num_redundant_experts_per_rank": 2, + "layers": { + "3": [[0, 1, 2, 3], [2, 3, 4, 5], [4, 5, 6, 7], [6, 7, 0, 1]] + } +} +``` + +写入前会再次校验全部层。实现使用独占创建的 `.lock` 文件避免多个服务同时写同一路径,并在成功写入后清除读取缓存。保存失败只记录 warning,不会中断在线推理。 + +## 10. 指标与运行行为 + +rank 0 周期性上报: + +```text +lightllm_eplb_topk_expert_imbalance_ratio +``` + +该指标先计算每层 `max(expert_load) / mean(expert_load)`,再对有效层求平均。值越接近 1,表示观测窗口内的逻辑专家负载越均衡。 + +`--eplb_rebalance_count` 的行为如下: + +- `-1`:持续允许动态规划和重排; +- `0`:使用初始或配置文件布局,不进行动态重排; +- 正整数:只统计实际完成且发生布局变化的重排;样本不足和布局不变不计数。 + +达到次数上限后,manager 仍会周期性采集和上报负载指标,但不再执行全局负载汇总和布局规划。 + +## 11. 当前限制 + +启用 EPLB 时需要满足: + +- 同时设置 `--enable_ep_moe`; +- `world_size > 1`; +- 逻辑专家数可以被 world size 整除; +- 冗余专家数大于 0,且不能超过本 rank 之外可复制的逻辑专家数; +- 同一 EP 通信组使用一致的 EPLB 参数; +- 不能与 `--enable_prefill_cudagraph` 同时使用; +- 当前不支持 `--enable_rl` 组合; +- 当前不支持 SM100 GPU; +- 当前动态 EPLB 执行路径依赖 EP DeepGEMM MoE 实现。 + +这些限制会在参数校验或 EPLB 初始化阶段尽早失败,避免服务带着不一致的布局进入推理。 + +## 12. 扩展新的规划算法 + +新增 planner 时建议遵循以下步骤: + +1. 在 `eplb/placement/` 中实现 `EPLBPlanner`; +2. 接收逻辑专家负载和当前完整布局,返回同 shape 的合法目标布局; +3. 在 `placement/factory.py` 的 builder 表中注册新的 `plan_mode`; +4. 在 CLI 和 `StartArgs` 的 `eplb_plan_mode` choices 中加入新名称; +5. 更新中英文参数文档; +6. 增加布局合法性、确定性、热点负载和迁移量测试; +7. 分别评估 prefill、decode 和混合流量,不要假设一种算法对所有部署形态都最优。 + +规划器必须保持以下边界: + +- 不直接修改 GPU 权重或路由 metadata; +- 不执行分布式通信; +- 不改变每个 rank 的物理槽位数; +- 不遗漏逻辑专家,不在同一 rank 重复放置同一专家; +- 相同输入返回确定结果; +- 尽量复用当前 rank 和物理槽位,避免均衡收益被迁移成本抵消。 + +`EPLBManager`、传输规划器和提交逻辑只依赖抽象的完整布局,因此新增算法不需要修改状态机。 + +## 13. 测试覆盖 + +EPLB 单元测试主要位于 `unit_tests/common/fused_moe/test_eplb.py`,覆盖: + +- 初始布局、路由 metadata 和拓扑排序; +- planner 输入校验、布局合法性、负载均衡和槽位复用; +- 状态机各分支及后台任务轮询; +- 链式依赖、环形依赖和并发传输批次; +- 权重与 metadata 的安全提交; +- 配置文件加载、校验和持久化; +- logical-to-physical kernel 的精确映射; +- token index `0..4096`、expert ID `0..255`,以及 `2、3、4、5、101、127、128、251` 个副本时的哈希分布。 + +多 GPU pinned-memory 传输测试位于 `unit_tests/common/fused_moe/test_eplb_transfer_gpu.py`。 diff --git a/docs/CN/source/index.rst b/docs/CN/source/index.rst index 8f79e5126f..2ca6e962c7 100755 --- a/docs/CN/source/index.rst +++ b/docs/CN/source/index.rst @@ -80,6 +80,7 @@ Lightllm 整合了众多的开源方案的优点,包括但不限于 FasterTran 架构介绍 token attention介绍 峰值显存调度器介绍 + EPLB 专家负载均衡实现 .. Indices and tables .. ================== From 3ac7954de9daf219d81aa64fbee40568306c4252 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 22 Sep 2026 08:43:36 +0000 Subject: [PATCH 64/72] fix --- .../meta_weights/fused_moe/impl/deepgemm_impl.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 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 c64002ed7b..f3bc3997f0 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 @@ -1,11 +1,6 @@ import torch from typing import Optional, Tuple, Any from .base_impl import FuseMoeBaseImpl -from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( - build_initial_local_expert_ids, - build_logical_to_physical_map, - load_layer_placement, -) from lightllm.distributed import dist_group_manager from lightllm.common.quantization.quantize_method import WeightPack from lightllm.utils.envs_utils import ( @@ -51,6 +46,13 @@ def _init_eplb_runtime(self): self.num_redundant_experts_per_rank = start_args.eplb_num_redundant_experts_per_rank if self.num_redundant_experts_per_rank > 0: + # 延迟导入:顶层导入会经 mode_backend 包形成 meta_weights -> server 的循环依赖。 + from lightllm.server.router.model_infer.mode_backend.eplb.placement import ( + build_initial_local_expert_ids, + build_logical_to_physical_map, + load_layer_placement, + ) + self.num_total_physical_experts = self.n_routed_experts + world_size * self.num_redundant_experts_per_rank # 阶段 1:先构造确定性的默认布局。未指定配置文件,或配置读取、校验失败时, From eb61641cd31d2a6295a32cde4fcf82b00054250b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 22 Sep 2026 10:16:58 +0000 Subject: [PATCH 65/72] fix --- .../layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py | 2 +- .../router/model_infer/mode_backend/eplb/runtime_manager.py | 2 +- 2 files changed, 2 insertions(+), 2 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 f3bc3997f0..bc0ee4042a 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 @@ -151,7 +151,7 @@ def _prepare_expert_execution( logical_to_physical_map=self.logical_to_physical_map, logical_expert_counter=self.route_counter, update_logical_expert_counter=self.recording, - mode="current_gpu_first", + mode="global_first", ) return topk_weights, topk_ids diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index fb964d2409..90ffcfcedd 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -36,7 +36,7 @@ logger = init_logger(__name__) EPLB_EXPERT_ALIGNMENT = 128 -EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT = 256 +EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT = 128 EPLB_EXPERT_IMBALANCE_RATIO_METRIC = "lightllm_eplb_topk_expert_imbalance_ratio" From cb1085073f9cd79995827f06023960c588c558f9 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 24 Sep 2026 08:01:24 +0000 Subject: [PATCH 66/72] fix --- docs/CN/source/framework/eplb.md | 71 ++++++++++++++++++++++++++------ 1 file changed, 59 insertions(+), 12 deletions(-) diff --git a/docs/CN/source/framework/eplb.md b/docs/CN/source/framework/eplb.md index 578150f095..7e36a2d30c 100644 --- a/docs/CN/source/framework/eplb.md +++ b/docs/CN/source/framework/eplb.md @@ -244,9 +244,56 @@ WAIT_PLAN_PLACEMENT_FINISHED 只有当整个 world 的平均样本量达到每个“层 × 逻辑专家”256 个 token 时才开始规划。样本不足不会清空 counter,低流量服务可以跨多个评估周期累计数据。 -## 7. 布局规划 +## 7. 专家分布分析 -### 7.1 规划器接口与选择 +基于 DeepSeek-R1(EP8 + DP8,单机 8xH200)加载 ShareGPT 语料(6000 条,input 512-2048 token) +实测的路由负载分布,用于校准布局规划(第 8 节)的设计假设。 + +### 7.1 各 rank 分布相似性(实测) + +在 `EPLBManager` 全局汇总负载(`_step_plan_placement` 的 `all_gather` 之后)时, +把各 rank 的 `[layer][logical_expert]` 负载逐层归一化为概率分布,以 rank0 的分布为基准, +与其余 rank 逐层计算 cosine 相似度。观测窗口为 warmup 阶段一次完整采样 +(58 个 MoE 层 × 7 对,共 406 对): + +| 指标 | 数值 | +| --- | --- | +| cosine 全体 min / 中位 / max | 0.9432 / 0.9728 / 0.9972 | +| 各 rank 对 rank0 的均值 | 0.971 ~ 0.977 | +| 最不相似层(层均值) | layer 43 (0.954)、32 (0.958)、41 (0.959) | +| 最相似层(层均值) | layer 0 (0.997)、1 (0.994)、4 (0.994) | + +探针日志示例: + +```text +eplb load prob cosine layer=0 vs_rank1..7: 0.9967 0.9965 0.9969 0.9962 0.9961 0.9972 0.9963 +eplb load prob cosine summary mean_by_rank(0..7): 1.0000 0.9773 0.9711 0.9745 0.9710 0.9762 0.9749 0.9735 +``` + +结论:**所有层上各 rank 的专家负载分布形状高度一致**。0.94~0.99 之间的小幅差异 +主要来自每 rank 仅承载约 1/8 流量的采样噪声,而非分布本身存在 rank 间异质性。 +DP 随机分流下每个 rank 的路由统计都是对全局路由分布的无偏采样, +分布形状(倾斜度、热点名单、长尾形态)在 rank 之间同源。 + +### 7.2 设计依据:分布规划只使用 rank0 的数据 + +上述相似性是后续设计中**分布规划只使用 rank0 负载统计**的合理性依据: + +1. **代表性**:各 rank 分布形状同源且高度一致(cosine 中位 0.97,浅层几乎重合), + rank0 的逐层概率分布与全局聚合分布在形状上等价, + 以 rank0 为样本做倾斜度分档、热点估计、副本预算分配(注水)不会产生系统性偏差。 +2. **成本**:规划无需等待全 rank 负载汇聚即可获得分布形状, + 采集与规划的关键路径缩短到单 rank 统计,状态机的全局同步点相应减少。 +3. **误差边界**:单 rank 估计相对全局的偏差上界为观测到的采样噪声 + (cosine ≥ 0.94),配合负载估计的平滑处理与迁移增益门槛, + 不会因单 rank 采样波动触发错误的副本迁移决策。 + +因此布局规划中所有"分布形状"相关的决策(概率分布、倾斜度、热点排序) +均以 rank0 的 logical expert 负载统计为准。 + +## 8. 布局规划 + +### 8.1 规划器接口与选择 所有布局算法实现统一的 `EPLBPlanner.plan(logical_expert_load, current_placement)` 接口,返回: @@ -258,7 +305,7 @@ WAIT_PLAN_PLACEMENT_FINISHED 规划只在 rank 0 的后台线程执行。完成后,目标布局通过控制通信组广播给所有 rank。相同输入必须产生确定结果,便于所有 rank 生成一致的传输计划。 -### 7.2 Greedy 规划算法 +### 8.2 Greedy 规划算法 默认算法按层独立规划,主要步骤如下: @@ -270,9 +317,9 @@ WAIT_PLAN_PLACEMENT_FINISHED 规划负载按 128 token 对齐,降低很小的计数波动对布局的影响。专家 ID 和 rank ID 用作稳定的平局规则,因此结果是确定性的。 -## 8. 权重迁移与安全提交 +## 9. 权重迁移与安全提交 -### 8.1 传输计划 +### 9.1 传输计划 传输规划器逐层比较当前布局和目标布局,为每个变化的目标槽绑定一个确定的源槽。选择源槽时优先使用不会被覆盖的稳定副本;没有稳定副本时,循环使用当前已有副本,避免把读取集中在同一个 rank。 @@ -284,7 +331,7 @@ WAIT_PLAN_PLACEMENT_FINISHED 每层单独生成批次,再按层顺序拼接。这样完成一层的提交后就能立即发布该层的新路由 metadata。 -### 8.2 数据路径 +### 9.2 数据路径 远程专家的传输路径为: @@ -300,7 +347,7 @@ WAIT_PLAN_PLACEMENT_FINISHED 控制面和权重传输分别使用独立的 Gloo process group,避免两类通信相互干扰。后台线程只负责把数据传入 pinned memory,不直接修改 live 权重。 -### 8.3 提交边界 +### 9.3 提交边界 只有当所有 rank 都确认当前批次传输完成后,主推理线程才会在 overlap stream 上统一: @@ -313,7 +360,7 @@ WAIT_PLAN_PLACEMENT_FINISHED 全部批次完成后,manager 发布目标布局、清空 route counter、增加完成次数,并回到 `COLLECTING`。 -## 9. 布局持久化 +## 10. 布局持久化 成功完成重排后,rank 0 会把最新完整布局写回 `--eplb_config_path`。配置内容包括: @@ -331,7 +378,7 @@ WAIT_PLAN_PLACEMENT_FINISHED 写入前会再次校验全部层。实现使用独占创建的 `.lock` 文件避免多个服务同时写同一路径,并在成功写入后清除读取缓存。保存失败只记录 warning,不会中断在线推理。 -## 10. 指标与运行行为 +## 11. 指标与运行行为 rank 0 周期性上报: @@ -349,7 +396,7 @@ lightllm_eplb_topk_expert_imbalance_ratio 达到次数上限后,manager 仍会周期性采集和上报负载指标,但不再执行全局负载汇总和布局规划。 -## 11. 当前限制 +## 12. 当前限制 启用 EPLB 时需要满足: @@ -365,7 +412,7 @@ lightllm_eplb_topk_expert_imbalance_ratio 这些限制会在参数校验或 EPLB 初始化阶段尽早失败,避免服务带着不一致的布局进入推理。 -## 12. 扩展新的规划算法 +## 13. 扩展新的规划算法 新增 planner 时建议遵循以下步骤: @@ -388,7 +435,7 @@ lightllm_eplb_topk_expert_imbalance_ratio `EPLBManager`、传输规划器和提交逻辑只依赖抽象的完整布局,因此新增算法不需要修改状态机。 -## 13. 测试覆盖 +## 14. 测试覆盖 EPLB 单元测试主要位于 `unit_tests/common/fused_moe/test_eplb.py`,覆盖: From 849357137b01636e569a375b32e549ebfc85ac26 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sat, 26 Sep 2026 07:15:45 +0000 Subject: [PATCH 67/72] refactor(eplb): record per-prefill route samples --- docs/CN/source/framework/eplb.md | 85 +++++-- .../fused_moe/fused_moe_weight.py | 3 +- .../fused_moe/gpt_oss_fused_moe_weight_tp.py | 3 +- .../meta_weights/fused_moe/impl/base_impl.py | 10 +- .../fused_moe/impl/deepgemm_impl.py | 25 +- .../fused_moe/impl/marlin_impl.py | 2 +- .../fused_moe/impl/triton_impl.py | 4 +- .../triton_kernel/fused_moe/eplb_topk_ids.py | 188 ++++++++++++--- .../layer_infer/transformer_layer_infer.py | 1 + .../layer_infer/transformer_layer_infer.py | 1 + .../layer_infer/transformer_layer_infer.py | 1 + .../layer_infer/transformer_layer_infer.py | 1 + .../layer_infer/transformer_layer_infer.py | 1 + .../mode_backend/eplb/runtime_manager.py | 53 +++-- unit_tests/common/fused_moe/test_eplb.py | 219 +++++++++++++----- 15 files changed, 458 insertions(+), 139 deletions(-) diff --git a/docs/CN/source/framework/eplb.md b/docs/CN/source/framework/eplb.md index 7e36a2d30c..052726a8c5 100644 --- a/docs/CN/source/framework/eplb.md +++ b/docs/CN/source/framework/eplb.md @@ -64,7 +64,7 @@ python -m lightllm.server.api_server \ │ logical top-k IDs v EPLB 路由 kernel - ├── 按 logical expert 累计 route_counter + ├── 把本次 prefill 的 logical expert 负载写入环形采样行 ├── 查询 logical_to_physical_map └── 为每个 token 选择 physical expert ID │ @@ -73,7 +73,7 @@ EPLB 路由 kernel 周期性控制面: -route_counter +prefill_route_counter(最近 24 次 prefill 采样) -> 全局负载汇总 -> placement planner -> target placement @@ -131,7 +131,8 @@ rank 3: [6, 7, 0, 1] - `local_logics_expert_ids_list`:本 rank 每个物理槽对应的逻辑专家; - `logical_to_physical_map`:logical ID 到可用 physical IDs 的设备路由表; -- `route_counter`:长度为 `E` 的 `int64` GPU 计数器; +- `prefill_route_counter`:shape 为 `[24, E]` 的 `int64` GPU 环形采样缓冲区; +- `prefill_route_sample_index`:shape 为 `[2]` 的 `int64` GPU 状态,分别保存 sample index 和核内同步计数; - `num_redundant_experts_per_rank`:本 rank 的额外槽位数。 目前启用 EP MoE 时使用 `FuseMoeDeepGEMM` 实现。EPLB manager 只收集启用了 EP 的 `layer.experts`,并保留模型中的层顺序。 @@ -180,7 +181,7 @@ _select_experts | `current_node_first` | 本节点存在副本时在节点内选择,否则回退到全局副本 | | `global_first` | 直接在全局全部有效副本间选择 | -当前 DeepGEMM EPLB 路径使用 `current_gpu_first`。未来如果要支持“本卡 -> 本节点 -> 全局”的三级回退,需要布局规划算法同时具备节点拓扑感知能力。 +当前 DeepGEMM EPLB 路径使用 `global_first`。未来如果要支持“本卡 -> 本节点 -> 全局”的三级回退,需要布局规划算法同时具备节点拓扑感知能力。 ### 5.4 副本哈希 @@ -188,11 +189,67 @@ _select_experts 该变换由奇数乘法和可逆的异或移位组成,可以显著降低规律性 token 间隔与副本数之间的低位相关性。例如同一专家每隔 4 个 token 出现且有 4 个副本时,简单线性哈希可能退化到单一副本,avalanche mix 能将流量重新打散。 -### 5.5 负载统计 +### 5.5 prefill 环形采样 -路由 kernel 在 physical ID 映射前,按 logical expert 对 `route_counter` 执行原子累加。这样同一逻辑专家的多个物理副本不会拆散规划器观察到的负载信号。 +负载采样只在调用方明确传入 `is_prefill=True` 时启用。decode 仍然执行 logical-to-physical 映射,但不会更新采样缓冲区,也不会推进 sample index。这样可以只使用吞吐量较大、统计稳定性更好的 prefill 路由结果,同时避免给高频 decode 路径增加原子操作。 -计数器与 MoE forward 位于同一条 overlap stream。清零操作也提交到该 stream,从而自然排在此前 forward 之后、后续 forward 之前,无需额外的全设备同步。 +每层使用一个 `[24, E]` 的 `prefill_route_counter`。其中每一行表示一次 prefill 路由 kernel 调用的 logical expert 直方图,24 表示最多保留最近 24 次采样,而不是 24 个 token 或 24 个 manager step。一次 manager step 内如果发生多次 prefill dispatch,它们会分别占用不同的采样行;第 25 次采样开始按环形方式覆盖最旧的数据: + +```text +sample_row = sample_index % 24 + +prefill_route_counter + row 0 -> 一次完整 prefill dispatch 的 [expert_0, ..., expert_E-1] 计数 + row 1 -> 下一次完整 prefill dispatch 的计数 + ... + row 23 -> 最近 24 次采样中的一行 +``` + +固定 24 行可以限制设备内存和 CPU 快照成本,并让规划器观察最近一段时间的流量,而不是让很早以前的流量永久影响当前布局。当前 manager 在评估时沿 sample 维求和,将 `[24, E]` 聚合回 `[E]`,因此现有 planner 接口无需感知环形缓冲区。 + +计数始终使用 logical expert ID,而不是最终选中的 physical expert ID。同一逻辑专家即使拥有多个物理副本,规划器看到的仍然是一份完整需求量,不会因为副本分发而被拆散。 + +### 5.6 无额外清零 kernel 的采样事务 + +环形行在复用前必须清零,否则新旧两次采样会叠加。为避免每次 prefill 额外发射一个清零 kernel,清零、路由计数和 sample index 提交都融合在 `_eplb_repair_topk_ids_kernel` 内;其中 `_record_prefill_route_sample` 是 Triton 子 JIT 函数,不会形成独立的 kernel launch。 + +`prefill_route_sample_index` 的两个元素含义如下: + +```text +[0] sample index:单调递增;对 24 取余得到当前环形行 +[1] sync state :0 表示目标行尚未清零 + 1 表示清零完成,采样可以开始 + 1 + completed_programs 表示已经完成的 program 数 +``` + +一次采样事务按以下顺序执行: + +1. 所有 program 读取同一个 sample index,并计算本次目标行。sample index 只由最后完成者推进,因此在本次 kernel 生命周期内保持不变。 +2. `program_id == 0` 清空目标行,然后通过带 `release` 语义的原子加一把 sync state 从 0 发布为 1。 +3. 其他 program 使用带 `acquire` 语义的原子读等待 sync state 达到 1,确保不会在清零完成前向目标行累加。 +4. 屏障通过后,每个 program 根据自己处理的有效 top-k 元素,对对应 logical expert 执行 `atomic_add(1)`。 +5. 每个 program 完成 physical ID 写回和负载计数后,再对 sync state 原子加一,提交一个完成信号。 +6. Triton 的 `atomic_add` 返回加法前的旧值,因此用 `old_value + 1 == num_programs + 1` 判断唯一的最后完成者。额外的 1 是步骤 2 发布的 ready 标记。 +7. 最后完成者先把 sample index 加一,再把 sync state 交换为 0,使下一次 kernel 可以复用后续环形行。 + +完整状态变化如下: + +```text +sync=0 + -> program 0 清零目标行 + -> sync=1(ready) + -> 所有 program 统计 logical expert 并分别提交完成信号 + -> sync=1+num_programs + -> 唯一最后完成者推进 sample index,并复位 sync=0 +``` + +ready 发布使用 `release`、等待方使用 `acquire`,最后完成信号使用 `acq_rel`,从而约束目标行清零和后续原子计数的可见顺序。所有调用还必须在同一 CUDA stream 上串行复用同一组 counter 和同步状态;当前 MoE forward 与采样都位于 overlap stream,满足这一约束。 + +### 5.7 manager 聚合与重置 + +manager 在安全推理边界把各层 `[24, E]` 环形缓冲区沿第 0 维求和并复制到 CPU,得到 planner 使用的 `[layer, logical_expert]` 负载。如果样本量不足或规划结果未改变布局,不主动清空缓冲区;后续 prefill 会继续写入,并在容量用满后滚动覆盖最旧行。 + +初始化、成功切换到新布局,以及达到重排次数上限后开始下一轮指标窗口时,manager 会同时清零 `prefill_route_counter` 和 `prefill_route_sample_index`。清零提交到 overlap stream,自然排在此前 forward 之后、后续 forward 之前,不需要额外的全设备同步。 ## 6. EPLB 状态机 @@ -200,7 +257,7 @@ _select_experts ```text [COLLECTING] - 累计 logical route counter,等待评估周期 + 将 prefill logical route 写入 24 行环形采样,等待评估周期 | v [EVALUATING] @@ -233,16 +290,16 @@ _select_experts ```text EVALUATING - |-- 平均 token 数不足 ----------> 保留 counter,继续累计 - `-- 达到重排次数上限 ----------> 清空 counter,只做周期性指标上报 + |-- 平均 token 数不足 ----------> 保留环形窗口,继续滚动采样 + `-- 达到重排次数上限 ----------> 清空采样,只做周期性指标上报 WAIT_PLAN_PLACEMENT_FINISHED - `-- 目标布局与当前布局相同 -----> 保留 counter,等待下次评估 + `-- 目标布局与当前布局相同 -----> 保留环形窗口,等待下次评估 ``` 默认每 20 个采样 step 评估一次,可以通过环境变量 `LIGHTLLM_EPLB_STEP_INTERVAL` 调整。该值必须大于 0。 -只有当整个 world 的平均样本量达到每个“层 × 逻辑专家”256 个 token 时才开始规划。样本不足不会清空 counter,低流量服务可以跨多个评估周期累计数据。 +只有当整个 world 的平均样本量达到每个“层 × 逻辑专家”128 个 token 时才开始规划。样本不足不会清空环形缓冲区,低流量服务可以跨多个评估周期继续采样;缓冲区写满后只保留最近 24 次 prefill dispatch。 ## 7. 专家分布分析 @@ -358,7 +415,7 @@ DP 随机分流下每个 rank 的路由统计都是对全局路由分布的无 权重和路由 metadata 在同一条 stream 上更新,后续 forward 只能看到完整提交后的状态,不会观察到“新路由指向旧权重”或“旧路由指向新权重”的中间状态。 -全部批次完成后,manager 发布目标布局、清空 route counter、增加完成次数,并回到 `COLLECTING`。 +全部批次完成后,manager 发布目标布局、清空 prefill 路由采样及设备端 sample index、增加完成次数,并回到 `COLLECTING`。 ## 10. 布局持久化 @@ -446,6 +503,8 @@ EPLB 单元测试主要位于 `unit_tests/common/fused_moe/test_eplb.py`,覆 - 权重与 metadata 的安全提交; - 配置文件加载、校验和持久化; - logical-to-physical kernel 的精确映射; +- prefill 环形采样的逐行计数、循环覆盖、sample index 推进和同步状态复位; +- 非 prefill 路径不修改采样状态,以及空输入不推进 sample index; - token index `0..4096`、expert ID `0..255`,以及 `2、3、4、5、101、127、128、251` 个副本时的哈希分布。 多 GPU pinned-memory 传输测试位于 `unit_tests/common/fused_moe/test_eplb_transfer_gpu.py`。 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 c04a4022dc..82546acf2e 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 @@ -109,10 +109,11 @@ def experts( use_grouped_topk: bool, topk_group: int, num_expert_group: int, - is_prefill: Optional[bool] = None, + is_prefill: bool, infer_state=None, shared_expert_gate: Optional[torch.Tensor] = None, ) -> torch.Tensor: + assert is_prefill is not None, "is_prefill must be explicitly specified for fused MoE execution" # Captures MoE topk expert ids for routed-experts metadata when enabled. moe_capture_callback = get_moe_capture_callback(infer_state, self.layer_num_) return self.fuse_moe_impl( diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/gpt_oss_fused_moe_weight_tp.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/gpt_oss_fused_moe_weight_tp.py index 240bc726ca..8bc37d74c7 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/gpt_oss_fused_moe_weight_tp.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/gpt_oss_fused_moe_weight_tp.py @@ -144,10 +144,11 @@ def experts( use_grouped_topk: bool, topk_group: int, num_expert_group: int, - is_prefill: Optional[bool] = None, + is_prefill: bool, infer_state=None, shared_expert_gate: Optional[torch.Tensor] = None, ): + assert is_prefill is not None, "is_prefill must be explicitly specified for fused MoE execution" assert shared_expert_gate is None, "shared_expert_gate is not supported by GPT-OSS fused MoE" topk_weights, topk_ids = self._router(router_logits, top_k) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py index ca42db39cf..dc6660925c 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py @@ -61,7 +61,7 @@ def __call__( use_grouped_topk: bool, topk_group: int, num_expert_group: int, - is_prefill: Optional[bool] = None, + is_prefill: bool, # Callback to capture MoE topk expert ids (routed experts metadata). moe_capture_callback: Optional[Callable[[torch.Tensor], None]] = None, per_expert_scale: Optional[torch.Tensor] = None, @@ -70,6 +70,7 @@ def __call__( # 追加 shared expert 时使用 sigmoid(logit) 作为其聚合权重。 shared_expert_gate: Optional[torch.Tensor] = None, ) -> torch.Tensor: + assert is_prefill is not None, "is_prefill must be explicitly specified for fused MoE execution" topk_weights, topk_ids = self._select_experts( input_tensor=input_tensor, router_logits=router_logits, @@ -88,6 +89,7 @@ def __call__( topk_weights=topk_weights, topk_ids=topk_ids, shared_expert_gate=shared_expert_gate, + is_prefill=is_prefill, ) return self._fused_experts( input_tensor=input_tensor, @@ -127,6 +129,7 @@ def _prepare_expert_execution( self, topk_weights: torch.Tensor, topk_ids: torch.Tensor, + is_prefill: bool, shared_expert_gate: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: """将逻辑路由结果转换为 MoE kernel 实际需要的执行布局。 @@ -143,6 +146,9 @@ def _prepare_expert_execution( 的权重,使其输出按 token 动态参与 routed expert 输出的聚合;传入 ``None`` 时,普通 fused shared expert 的追加权重为 1。EP 路径不使用该参数,而是 单独计算 shared expert,并在应用相同门控后与 routed MoE 输出相加。 + + ``is_prefill`` 必须由上层入口显式传入 ``True`` 或 ``False``。即使当前实现 + 尚未使用该信息,也不允许用 ``None`` 隐式表示执行阶段。 """ pass @@ -154,8 +160,8 @@ def _fused_experts( w2: WeightPack, topk_weights: torch.Tensor, topk_ids: torch.Tensor, + is_prefill: bool, router_logits: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, ) -> torch.Tensor: """根据准备完成的路由结果执行融合 MoE 计算。 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 bc0ee4042a..82b7309967 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 @@ -89,8 +89,16 @@ def _init_eplb_runtime(self): ), dtype=torch.int32, ).cuda() - # 始终按逻辑专家统计负载,冗余副本不会拆散规划器观察到的负载信号。 - self.route_counter = torch.zeros(self.n_routed_experts, dtype=torch.int64, device="cuda") + # 环形缓冲区保留最近 24 次 prefill 路由采样,每次采样写入独立的一行; + # 始终按 logical expert 统计,冗余副本不会拆散规划器观察到的负载信号。 + self.prefill_route_counter = torch.zeros( + (24, self.n_routed_experts), + dtype=torch.int64, + device="cuda", + ) + # [0] 是单调递增的 sample index;[1] 用于在同一个 kernel 内协调 + # 目标行清零,并从所有 program 中选出最后完成者。 + self.prefill_route_sample_index = torch.zeros(2, dtype=torch.int64, device="cuda") # 动态 EPLB 默认采集路由负载;以后使用配置文件固定专家布局时, # 可以关闭该开关,避免执行不再需要的 atomic counter 更新。 self.recording = True @@ -142,15 +150,18 @@ def _prepare_expert_execution( self, topk_weights: torch.Tensor, topk_ids: torch.Tensor, + is_prefill: bool, shared_expert_gate: Optional[torch.Tensor] = None, ): + assert is_prefill is not None, "is_prefill must be explicitly specified for fused MoE execution" assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" if self.num_redundant_experts_per_rank > 0: topk_ids = eplb_repair_topk_ids( logical_topk_ids=topk_ids, logical_to_physical_map=self.logical_to_physical_map, - logical_expert_counter=self.route_counter, - update_logical_expert_counter=self.recording, + prefill_route_counter=self.prefill_route_counter, + prefill_route_sample_index=self.prefill_route_sample_index, + update_prefill_route_counter=self.recording and is_prefill is True, mode="global_first", ) return topk_weights, topk_ids @@ -162,8 +173,8 @@ def _fused_experts( w2: WeightPack, topk_weights: torch.Tensor, topk_ids: torch.Tensor, + is_prefill: bool, router_logits: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, ): output = fused_experts( hidden_states=input_tensor, @@ -201,7 +212,7 @@ def low_latency_dispatch( num_expert_group=n_group, scoring_func=scoring_func, ) - topk_weights, topk_idx = self._prepare_expert_execution(topk_weights, topk_idx) + topk_weights, topk_idx = self._prepare_expert_execution(topk_weights, topk_idx, is_prefill=False) topk_idx = topk_idx.to(torch.long) num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_decode() @@ -241,7 +252,7 @@ def select_experts_and_quant_input( num_expert_group=n_group, scoring_func=scoring_func, ) - topk_weights, topk_idx = self._prepare_expert_execution(topk_weights, topk_idx) + topk_weights, topk_idx = self._prepare_expert_execution(topk_weights, topk_idx, is_prefill=True) qinput_tensor = quantize_fused_experts_input(hidden_states, w13, self.quant_method) return topk_weights, topk_idx.to(torch.long), qinput_tensor diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py index 2ee57fe916..c4937c681d 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py @@ -32,8 +32,8 @@ def _fused_experts( w2: WeightPack, topk_weights: torch.Tensor, topk_ids: torch.Tensor, + is_prefill: bool, router_logits: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, ): w1_weight, w1_scale, w1_zero_point = w13.weight, w13.weight_scale, w13.weight_zero_point diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index a42f9c9f36..886217420e 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -42,8 +42,10 @@ def _prepare_expert_execution( self, topk_weights: torch.Tensor, topk_ids: torch.Tensor, + is_prefill: bool, shared_expert_gate: Optional[torch.Tensor] = None, ): + assert is_prefill is not None, "is_prefill must be explicitly specified for fused MoE execution" if self.num_fused_shared_experts > 0: from lightllm.common.basemodel.triton_kernel.fused_moe.append_shared_expert_topk import ( append_fused_shared_experts, @@ -65,8 +67,8 @@ def _fused_experts( w2: WeightPack, topk_weights: torch.Tensor, topk_ids: torch.Tensor, + is_prefill: bool, router_logits: Optional[torch.Tensor] = None, - is_prefill: bool = False, ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py index 029f2fb48e..cd328e6eaa 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/eplb_topk_ids.py @@ -18,6 +18,115 @@ def _replica_index(token_index, logical_expert_id, num_valid_replicas): return value % num_valid_replicas.to(tl.uint32) +@triton.jit +def _record_prefill_route_sample( + logical_expert_ids, + valid_mask, + prefill_route_counter_ptr, + prefill_route_counter_row_stride, + prefill_route_sample_index_ptr, + NUM_LOGICAL_EXPERTS: tl.constexpr, + PREFILL_ROUTE_COUNTER_CAPACITY: tl.constexpr, + COUNTER_BLOCK_SIZE: tl.constexpr, +): + """在主路由 kernel 内完成一次 prefill 采样事务。""" + # 同步槽 ``prefill_route_sample_index[1]`` 的状态机: + # + # 0 : 本次采样尚未清零; + # 1 : program 0 已清零,ready 标记已经发布; + # 1 + completed : ready 标记加上已经完成的 program 数; + # 1 + num_programs : 本次所有 program 均已完成。 + # + # 所有调用必须排在同一 CUDA stream 上,不能并发复用同一组 counter + # 和同步状态。 + program_id = tl.program_id(0) + + # 1. 所有 program 读取相同的 sample index,并通过取余定位本次写入的 + # 环形行。sample index 只有在全部 program 完成后才会推进,因此本次 + # kernel 生命周期内,各 program 看到的目标行保持不变。 + sample_index = tl.load(prefill_route_sample_index_ptr) + sample_row = sample_index % PREFILL_ROUTE_COUNTER_CAPACITY + + if program_id == 0: + # 2. program 0 负责初始化目标行。正常入口处同步槽必须为 0;清零 + # 完成后,通过 release 原子加一发布 ready=1,使此前的 store 对 + # 随后获得 ready 标记的其他 program 可见。 + sync_state = tl.atomic_add( + prefill_route_sample_index_ptr + 1, + 0, + sem="acquire", + scope="gpu", + ) + if sync_state == 0: + expert_offsets = tl.arange(0, COUNTER_BLOCK_SIZE) + tl.store( + prefill_route_counter_ptr + sample_row * prefill_route_counter_row_stride + expert_offsets, + 0, + mask=expert_offsets < NUM_LOGICAL_EXPERTS, + ) + tl.atomic_add( + prefill_route_sample_index_ptr + 1, + 1, + sem="release", + scope="gpu", + ) + else: + # 3. 其他 program 使用 acquire 原子读等待 ready 标记。只有观察到 + # sync >= 1 后才能离开循环,从而保证不会与 program 0 的清零 store + # 并发访问同一个 counter 行。 + sync_state = tl.atomic_add( + prefill_route_sample_index_ptr + 1, + 0, + sem="acquire", + scope="gpu", + ) + while sync_state < 1: + sync_state = tl.atomic_add( + prefill_route_sample_index_ptr + 1, + 0, + sem="acquire", + scope="gpu", + ) + + # 4. 清零屏障通过后,各 program 将自己的 logical expert 路由结果原子 + # 累加到同一采样行。这里统计 logical expert,冗余 physical 副本不会 + # 拆散规划器观察到的负载信号。 + tl.atomic_add( + prefill_route_counter_ptr + sample_row * prefill_route_counter_row_stride + logical_expert_ids, + 1, + mask=valid_mask, + sem="relaxed", + ) + + # 5. 本 program 完成计数后,向同步槽提交一个完成信号。调用发生在主 + # kernel 的 physical ID 写回之后,因此该信号同时表示两部分工作均完成。 + # atomic_add 返回旧值,故完成后的新值需要显式加一。当新值等于 + # ``num_programs + 1`` 时,ready 标记和全部 program 的完成信号均已到达。 + completed_programs = tl.atomic_add( + prefill_route_sample_index_ptr + 1, + 1, + sem="acq_rel", + scope="gpu", + ) + sync_after_completion = completed_programs + 1 + is_last_program = sync_after_completion == tl.num_programs(0) + 1 + if is_last_program: + # 最后完成者提交本次事务:先推进 sample index,使下一次采样指向 + # 后续环形行;再把同步槽复位为 0,供下一次 program 0 执行清零。 + tl.atomic_add( + prefill_route_sample_index_ptr, + 1, + sem="release", + scope="gpu", + ) + tl.atomic_xchg( + prefill_route_sample_index_ptr + 1, + 0, + sem="release", + scope="gpu", + ) + + @triton.jit def _eplb_repair_topk_ids_kernel( logical_topk_ids_ptr, @@ -26,28 +135,24 @@ def _eplb_repair_topk_ids_kernel( top_k, logical_to_physical_map_ptr, logical_to_physical_map_row_stride, - logical_expert_counter_ptr, + prefill_route_counter_ptr, + prefill_route_counter_row_stride, + prefill_route_sample_index_ptr, DISPATCH_MODE: tl.constexpr, - UPDATE_LOGICAL_EXPERT_COUNTER: tl.constexpr, + UPDATE_PREFILL_ROUTE_COUNTER: tl.constexpr, + NUM_LOGICAL_EXPERTS: tl.constexpr, + PREFILL_ROUTE_COUNTER_CAPACITY: tl.constexpr, + COUNTER_BLOCK_SIZE: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): # 阶段 1:将二维 [num_tokens, top_k] 路由结果展平后分块处理。 # topk_id_offsets 同时用于访问输入、输出,并可恢复它所属的 token 下标。 - topk_id_offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + program_id = tl.program_id(0) + topk_id_offsets = program_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) valid_mask = topk_id_offsets < num_topk_ids logical_expert_ids = tl.load(logical_topk_ids_ptr + topk_id_offsets, mask=valid_mask, other=0) - # 阶段 2:在 EPLB 采样窗口内,按 logical expert 统计本次路由负载。 - # 统计发生在 physical ID 修复之前,因此冗余副本不会拆散逻辑专家负载。 - if UPDATE_LOGICAL_EXPERT_COUNTER: - tl.atomic_add( - logical_expert_counter_ptr + logical_expert_ids, - 1, - mask=valid_mask, - sem="relaxed", - ) - - # 阶段 3:定位每个 logical expert 的打包映射行。固定头部的布局为: + # 阶段 2:定位每个 logical expert 的打包映射行。固定头部的布局为: # # [0] 所有 rank 上的有效副本总数 # [1] 当前节点上的有效副本数,包含本卡副本 @@ -65,7 +170,7 @@ def _eplb_repair_topk_ids_kernel( other=0, ) - # 阶段 4:根据调用方显式指定的分发模式选择参与 hash 的候选前缀。 + # 阶段 3:根据调用方显式指定的分发模式选择参与 hash 的候选前缀。 if DISPATCH_MODE == 0: # current_gpu_first: 本卡 -> 全局。 # TODO: 等 EPLB 布局算法支持节点拓扑感知后,再考虑增加 @@ -88,7 +193,7 @@ def _eplb_repair_topk_ids_kernel( token_indices = topk_id_offsets // top_k selected_replica_indices = _replica_index(token_indices, logical_expert_ids, num_preferred_replicas) - # 阶段 5:读取选中槽位的 physical expert ID 并写入新的输出 tensor。 + # 阶段 4:读取选中槽位的 physical expert ID 并写入新的输出 tensor。 # logical_topk_ids 只读,后续 callback 仍可安全观察原始逻辑路由结果。 physical_expert_ids = tl.load( logical_to_physical_map_ptr + map_row_offsets + 3 + selected_replica_indices, @@ -97,13 +202,28 @@ def _eplb_repair_topk_ids_kernel( ) tl.store(physical_topk_ids_ptr + topk_id_offsets, physical_expert_ids, mask=valid_mask) + # 阶段 5:记录本次 prefill 路由采样。子函数负责目标行清零、program 间 + # ready 同步、logical expert 计数,以及最后完成者对 sample index 的提交。 + if UPDATE_PREFILL_ROUTE_COUNTER: + _record_prefill_route_sample( + logical_expert_ids, + valid_mask, + prefill_route_counter_ptr, + prefill_route_counter_row_stride, + prefill_route_sample_index_ptr, + NUM_LOGICAL_EXPERTS, + PREFILL_ROUTE_COUNTER_CAPACITY, + COUNTER_BLOCK_SIZE, + ) + @torch.no_grad() def eplb_repair_topk_ids( logical_topk_ids: torch.Tensor, logical_to_physical_map: torch.Tensor, - logical_expert_counter: torch.Tensor, - update_logical_expert_counter: bool, + prefill_route_counter: torch.Tensor, + prefill_route_sample_index: torch.Tensor, + update_prefill_route_counter: bool, mode: str, ) -> torch.Tensor: """将 logical top-k ID 转换为当前 EPLB 布局中的 physical expert ID。 @@ -119,10 +239,15 @@ def eplb_repair_topk_ids( physical IDs 按本卡、本节点其他卡、其他节点排列;具体参与分发的 候选前缀由 ``mode`` 决定。 有效副本之后未使用的 padding 槽位为 -1,kernel 不会读取它们。 - logical_expert_counter: 每个 logical expert 的累计路由次数,shape 为 - ``[num_logical_experts]``。 - update_logical_expert_counter: 是否将本次 logical 路由结果累计到 - ``logical_expert_counter``。固定布局不需要动态重排时可以关闭。 + prefill_route_counter: prefill 路由采样的环形缓冲区,shape 为 + ``[sample_capacity, num_logical_experts]``。一次 kernel 调用只写 + ``prefill_route_sample_index[0] % sample_capacity`` 对应的一行。 + prefill_route_sample_index: shape 为 ``[2]`` 的设备端同步状态。第 0 项 + 是单调递增的 sample index;第 1 项用于核内清零与完成同步:0 + 表示尚未清零,正数为 ready 标记 1 加上已完成的 program 数量; + 达到 ``num_programs + 1`` 后,最后完成者推进 sample index 并清零。 + update_prefill_route_counter: 是否记录本次 prefill 路由采样并推进 sample + index。decode 或固定布局不需要采样时可以关闭。 mode: 必须显式指定的副本分发模式,不提供默认值: * ``current_gpu_first``:本卡优先,没有本卡副本时回退到全局; @@ -147,8 +272,14 @@ def eplb_repair_topk_ids( assert logical_to_physical_map.ndim == 2 assert logical_to_physical_map.shape[1] > 3 assert logical_to_physical_map.stride(1) == 1 - assert logical_expert_counter.ndim == 1 - assert logical_expert_counter.shape[0] == logical_to_physical_map.shape[0] + assert prefill_route_counter.ndim == 2 + assert prefill_route_counter.shape[0] > 1 + assert prefill_route_counter.shape[1] == logical_to_physical_map.shape[0] + assert prefill_route_counter.is_contiguous() + assert prefill_route_counter.dtype is torch.int64 + assert prefill_route_sample_index.shape == (2,) + assert prefill_route_sample_index.dtype is torch.int64 + assert prefill_route_sample_index.device == prefill_route_counter.device physical_topk_ids = torch.empty_like(logical_topk_ids) if logical_topk_ids.numel() == 0: return physical_topk_ids @@ -161,9 +292,14 @@ def eplb_repair_topk_ids( top_k=logical_topk_ids.shape[1], logical_to_physical_map_ptr=logical_to_physical_map, logical_to_physical_map_row_stride=logical_to_physical_map.stride(0), - logical_expert_counter_ptr=logical_expert_counter, + prefill_route_counter_ptr=prefill_route_counter, + prefill_route_counter_row_stride=prefill_route_counter.stride(0), + prefill_route_sample_index_ptr=prefill_route_sample_index, DISPATCH_MODE=dispatch_mode, - UPDATE_LOGICAL_EXPERT_COUNTER=update_logical_expert_counter, + UPDATE_PREFILL_ROUTE_COUNTER=update_prefill_route_counter, + NUM_LOGICAL_EXPERTS=prefill_route_counter.shape[1], + PREFILL_ROUTE_COUNTER_CAPACITY=prefill_route_counter.shape[0], + COUNTER_BLOCK_SIZE=triton.next_power_of_2(prefill_route_counter.shape[1]), BLOCK_SIZE=block_size, num_warps=4, num_stages=1, diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py index 3254031056..92ad86d21f 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -234,6 +234,7 @@ def _moe_ffn_tp( use_grouped_topk=self.n_group, topk_group=self.topk_group, num_expert_group=self.n_group, + is_prefill=infer_state.is_prefill, infer_state=infer_state, ) diff --git a/lightllm/models/gpt_oss/layer_infer/transformer_layer_infer.py b/lightllm/models/gpt_oss/layer_infer/transformer_layer_infer.py index 490d2dc4c5..1861b61848 100644 --- a/lightllm/models/gpt_oss/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/gpt_oss/layer_infer/transformer_layer_infer.py @@ -52,6 +52,7 @@ def _ffn(self, input, infer_state, layer_weight: GptOssTransformerLayerWeight) - use_grouped_topk=False, topk_group=None, num_expert_group=None, + is_prefill=infer_state.is_prefill, infer_state=infer_state, ) hidden_states = hidden_states.view(num_tokens, hidden_dim) diff --git a/lightllm/models/mixtral/layer_infer/transformer_layer_infer.py b/lightllm/models/mixtral/layer_infer/transformer_layer_infer.py index 8134dc266d..4e00659e80 100644 --- a/lightllm/models/mixtral/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/mixtral/layer_infer/transformer_layer_infer.py @@ -26,6 +26,7 @@ def _ffn(self, input, infer_state: InferStateInfo, layer_weight: MixtralTransfor use_grouped_topk=False, topk_group=None, num_expert_group=None, + is_prefill=infer_state.is_prefill, infer_state=infer_state, ) return hidden_states.view(num_tokens, hidden_dim) 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 7311c4d141..d33367ab28 100644 --- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py @@ -88,6 +88,7 @@ def _moe_ffn_tp( use_grouped_topk=False, topk_group=None, num_expert_group=None, + is_prefill=infer_state.is_prefill, infer_state=infer_state, ) return hidden_states.view(num_tokens, hidden_dim) diff --git a/lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py index 92d68c9fd2..1249b8ff34 100644 --- a/lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py @@ -95,6 +95,7 @@ def _moe_ffn_tp( use_grouped_topk=False, topk_group=None, num_expert_group=None, + is_prefill=infer_state.is_prefill, infer_state=infer_state, shared_expert_gate=shared_expert_gate, ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 90ffcfcedd..bf3953d331 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -58,13 +58,13 @@ class EPLBManager: 每次调用 :meth:`step` 最多处理一个状态。主路径及各状态的职责如下:: [COLLECTING] - 推理 kernel 持续累计各层 logical expert 的 route counter; + prefill kernel 将各层 logical expert 负载写入 24 行环形样本; manager 只记录采样 step,等待下一个评估周期。 | | 评估周期到达 v [EVALUATING] - 将本地 counter 快照到 CPU 并上报负载指标;汇总各 rank 的 + 将本地环形样本聚合到 CPU 并上报负载指标;汇总各 rank 的 token 总数,判断样本量和剩余重排次数。 | | 样本充足且仍允许重排 @@ -89,16 +89,16 @@ class EPLBManager: [TRANSFERRING] 分批启动并轮询后台权重传输;整批完成后,主推理线程在安全 边界统一提交权重和路由 metadata。全部批次完成后发布新布局、 - 清空 counter 并回到 COLLECTING。 + 清空 prefill 路由样本并回到 COLLECTING。 以下分支会提前回到 ``COLLECTING``:: EVALUATING - |-- 平均 token 数不足 --------> 保留 counter,继续累计样本 - `-- 已达到重排次数上限 ------> 清空 counter,仅周期性上报指标 + |-- 平均 token 数不足 --------> 保留环形窗口,继续滚动采样 + `-- 已达到重排次数上限 ------> 清空样本,仅周期性上报指标 WAIT_PLAN_PLACEMENT_FINISHED - `-- 目标布局与当前布局相同 ---> 保留 counter,等待下次评估 + `-- 目标布局与当前布局相同 ---> 保留环形窗口,等待下次评估 布局规划、传输规划和权重传输在后台执行;主推理线程只负责创建任务、 轮询状态,以及在安全边界提交已经完成的结果。 @@ -142,8 +142,8 @@ def __init__( expert_alignment=EPLB_EXPERT_ALIGNMENT, ) - # 评估调度:steps 只在 COLLECTING 状态递增。route counter 从当前 - # 布局生效时开始累计,让低流量服务可以跨多个评估周期收集足够样本。 + # 评估调度:steps 只在 COLLECTING 状态递增。prefill 路由样本从当前 + # 布局生效时开始写入,并在固定容量内保留最近的采样窗口。 self.step_interval: int = get_eplb_step_interval() self.steps: int = 0 self.max_rebalance_count: int = max_rebalance_count @@ -174,7 +174,7 @@ def __init__( self.state = EPLBManagerState.COLLECTING self.next_evaluation_step = self.step_interval - self._clear_route_counters() + self._clear_prefill_route_samples() if self.global_rank == 0: self.metric_client: MetricClient = MetricClient(get_shm_port_args().metric_port) @@ -230,24 +230,28 @@ def _step_collecting(self) -> None: def _step_evaluating(self) -> None: """发布本地负载指标,并在次数允许时根据全局样本量决定是否规划。""" - counters = [impl.route_counter for impl in self._eplb_impls] - if any(counter.ndim != 1 or counter.shape[0] != self.num_logical_experts for counter in counters): - raise RuntimeError("EPLB route counter shape must be [num_logical_experts]") - - # 将各层累计的路由计数复制到 CPU,后续规划统一使用这份快照。 - # 此处先不清零 GPU counter:如果样本不足或无需迁移,下一周期会 - # 继续累计;达到重排上限或成功切换到新布局后才重新开始统计。本轮 + counters = [impl.prefill_route_counter for impl in self._eplb_impls] + if any(counter.ndim != 2 or counter.shape[1] != self.num_logical_experts for counter in counters): + raise RuntimeError("EPLB prefill route counter shape must be [sample_capacity, num_logical_experts]") + if len({counter.shape[0] for counter in counters}) != 1: + raise RuntimeError("EPLB prefill route counter capacities must match across layers") + + # 将各层环形样本聚合并复制到 CPU,后续规划统一使用这份快照。 + # 此处先不清零 GPU 样本:如果样本不足或无需迁移,下一周期会继续 + # 滚动覆盖最旧行;达到重排上限或成功切换到新布局后才重置窗口。本轮 # 异步规划使用独立的 CPU 快照,不会与后续的 atomic add 竞争。 - local_load = torch.stack([counter.detach().cpu() for counter in counters]) + # 当前 planner 仍消费 [layer, logical_expert] 聚合负载;先对 24 行采样 + # 求和保持现有规划语义。后续逐样本规划可以直接在这里保留 sample 维。 + local_load = torch.stack([counter.detach().sum(dim=0).cpu() for counter in counters]) self._publish_expert_load_metric(local_load) # 达到重排次数上限后仍保留周期性负载上报,但不再执行后续的跨 rank - # 通信和布局规划。清空本轮计数,使下一次指标对应新的采样窗口。 + # 通信和布局规划。清空本轮样本,使下一次指标对应新的采样窗口。 reached_rebalance_limit = ( self.max_rebalance_count != -1 and self.completed_rebalance_count >= self.max_rebalance_count ) if reached_rebalance_limit: - self._clear_route_counters() + self._clear_prefill_route_samples() self.state = EPLBManagerState.COLLECTING else: # 汇集各 rank 的 token 总数,判断当前统计量是否足以进行布局规划。 @@ -381,7 +385,7 @@ def _step_transferring(self) -> None: elapsed = time.time() - self.rebalance_started_at if self.global_rank == 0: self._persist_current_placement() - self._clear_route_counters() + self._clear_prefill_route_samples() self.completed_rebalance_count += 1 del self.pending_transfer_batches del self.target_placement @@ -491,16 +495,17 @@ def _publish_expert_load_metric(self, local_load: torch.Tensor) -> None: _expert_load_imbalance_ratio(local_load), ) - def _clear_route_counters(self) -> None: - """在 overlap stream 上清空所有层的逻辑专家路由计数。""" + def _clear_prefill_route_samples(self) -> None: + """在 overlap stream 上清空所有层的 prefill 路由样本和设备端索引。""" from lightllm.server.router.model_infer.infer_batch import g_infer_context - # route counter 由 forward 中的 Triton kernel 在 overlap stream 上更新。 + # prefill route sample 由 forward 中的 Triton kernel 在 overlap stream 上更新。 # 将 zero_ 排到同一条 stream,可保证它位于此前 forward 之后、下一次 # forward 之前,无需额外 synchronize,也不会与 atomic add 并发。 with torch.cuda.stream(g_infer_context.get_overlap_stream()): for impl in self._eplb_impls: - impl.route_counter.zero_() + impl.prefill_route_counter.zero_() + impl.prefill_route_sample_index.zero_() def _persist_current_placement(self) -> None: """由 rank 0 将当前完整布局写回启动时指定的输入/输出文件。""" diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index cd85b8f1e3..596e4fdce8 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -58,13 +58,16 @@ def _test_moe_impl( num_logical_experts=128, world_size=16, num_redundant_experts_per_rank=1, - route_counter=None, + prefill_route_counter=None, + prefill_route_sample_index=None, recording=False, ): logical_to_physical_map = None if eplb: - if route_counter is None: - route_counter = torch.zeros((num_logical_experts,), dtype=torch.int64) + if prefill_route_counter is None: + prefill_route_counter = torch.zeros((24, num_logical_experts), dtype=torch.int64) + if prefill_route_sample_index is None: + prefill_route_sample_index = torch.zeros((2,), dtype=torch.int64) logical_to_physical_map = torch.zeros((num_logical_experts, world_size + 3), dtype=torch.int32) logical_to_physical_map[:, :3] = 1 else: @@ -75,7 +78,8 @@ def _test_moe_impl( num_redundant_experts_per_rank=num_redundant_experts_per_rank, local_logics_expert_ids_list=list(range(num_logical_experts // world_size + num_redundant_experts_per_rank)), logical_to_physical_map=logical_to_physical_map, - route_counter=route_counter, + prefill_route_counter=prefill_route_counter, + prefill_route_sample_index=prefill_route_sample_index, recording=recording, ) @@ -85,7 +89,8 @@ def _set_deepgemm_runtime(impl, runtime): "num_total_physical_experts", "num_redundant_experts_per_rank", "logical_to_physical_map", - "route_counter", + "prefill_route_counter", + "prefill_route_sample_index", "recording", ): setattr(impl, name, getattr(runtime, name)) @@ -134,8 +139,8 @@ def _select_experts( ): return "weights", "logical_ids" - def _prepare_expert_execution(self, topk_weights, topk_ids, shared_expert_gate=None): - seen["prepare"] = {"topk_ids": topk_ids} + def _prepare_expert_execution(self, topk_weights, topk_ids, is_prefill, shared_expert_gate=None): + seen["prepare"] = {"topk_ids": topk_ids, "is_prefill": is_prefill} return topk_weights, "physical_ids" def _fused_experts( @@ -145,8 +150,8 @@ def _fused_experts( w2, topk_weights, topk_ids, + is_prefill, router_logits=None, - is_prefill=None, ): seen["fused"] = {"topk_ids": topk_ids} return "output" @@ -166,12 +171,29 @@ def _fused_experts( 0, 0, moe_capture_callback=captured.append, + is_prefill=True, ) assert result == "output" assert captured == ["logical_ids"] assert seen["prepare"]["topk_ids"] == "logical_ids" + assert seen["prepare"]["is_prefill"] is True assert seen["fused"]["topk_ids"] == "physical_ids" + with pytest.raises(AssertionError, match="is_prefill must be explicitly specified"): + impl( + "input", + "logits", + "w13", + "w2", + None, + "softmax", + 2, + False, + False, + 0, + 0, + ) + def test_factory_selects_all_paths_without_ep_constructor_state(monkeypatch): plain_quant = SimpleNamespace(method_name="none") @@ -194,7 +216,7 @@ def test_factory_selects_all_paths_without_ep_constructor_state(monkeypatch): assert isinstance(ep_impl, deepgemm_module.FuseMoeDeepGEMM) assert ep_impl.num_total_physical_experts == 4 assert not hasattr(ep_impl, "num_primary_experts_per_rank") - assert not hasattr(ep_impl, "route_counter") + assert not hasattr(ep_impl, "prefill_route_counter") assert not hasattr(ep_impl, "expert_parallel_state") assert isinstance( create_fuse_moe_impl( @@ -760,15 +782,15 @@ def test_transfer_plan_respects_explicit_target_slots(): def test_manager_evaluating_copies_route_counters_to_cpu_without_modifying_them(monkeypatch): counters = [ - torch.tensor([10, 11], dtype=torch.int64), - torch.tensor([40, 41], dtype=torch.int64), + torch.tensor([[10, 11]], dtype=torch.int64), + torch.tensor([[40, 41]], dtype=torch.int64), ] manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.EVALUATING manager._eplb_impls = [ _test_moe_impl( eplb=True, - route_counter=counter, + prefill_route_counter=counter, num_logical_experts=2, world_size=1, ) @@ -791,8 +813,8 @@ def all_gather_object(output, local_token_count, **_kwargs): manager._step_evaluating() assert local_token_counts == [102] - assert torch.equal(counters[0], torch.tensor([10, 11], dtype=torch.int64)) - assert torch.equal(counters[1], torch.tensor([40, 41], dtype=torch.int64)) + assert torch.equal(counters[0], torch.tensor([[10, 11]], dtype=torch.int64)) + assert torch.equal(counters[1], torch.tensor([[40, 41]], dtype=torch.int64)) def test_manager_delegates_distribution_planning_to_planner_class(): @@ -921,7 +943,7 @@ def test_manager_does_not_publish_expert_load_metrics_from_other_ranks(): assert not hasattr(manager, "metric_client") -def test_eplb_route_counter_has_one_entry_per_logical_expert(monkeypatch): +def test_eplb_prefill_route_counter_has_24_samples_per_logical_expert(monkeypatch): args = type( "Args", (), @@ -945,7 +967,9 @@ def cpu_zeros(*shape, **kwargs): impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace()) - assert impl.route_counter.shape == (4,) + assert impl.prefill_route_counter.shape == (24, 4) + assert impl.prefill_route_sample_index.shape == (2,) + assert torch.equal(impl.prefill_route_sample_index, torch.zeros(2, dtype=torch.int64)) def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): @@ -966,7 +990,8 @@ def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): assert not hasattr(impl, "num_primary_experts_per_rank") assert not hasattr(impl, "initial_local_expert_ids_by_rank") assert not hasattr(impl, "logical_to_physical_map") - assert not hasattr(impl, "route_counter") + assert not hasattr(impl, "prefill_route_counter") + assert not hasattr(impl, "prefill_route_sample_index") assert not hasattr(impl, "recording") @@ -979,7 +1004,7 @@ def test_manager_evaluation_gathers_token_counts_from_all_ranks(monkeypatch): { "fuse_moe_impl": _test_moe_impl( eplb=True, - route_counter=torch.zeros((4,), dtype=torch.int64), + prefill_route_counter=torch.zeros((24, 4), dtype=torch.int64), num_logical_experts=4, world_size=1, ) @@ -996,8 +1021,9 @@ def test_manager_evaluation_gathers_token_counts_from_all_ranks(monkeypatch): manager.control_group = object() manager.max_rebalance_count = -1 manager.completed_rebalance_count = 0 - local = torch.full((4,), 100, dtype=torch.int64) - manager._eplb_impls[0].route_counter = local + local = torch.full((24, 4), 0, dtype=torch.int64) + local[0].fill_(100) + manager._eplb_impls[0].prefill_route_counter = local manager.state = manager_module.EPLBManagerState.EVALUATING seen = {} @@ -1131,8 +1157,8 @@ def repair(**kwargs): assert topk_idx.dtype is torch.long assert qinput == "qinput" assert calls[0]["logical_topk_ids"] is logical_ids - assert not calls[0]["update_logical_expert_counter"] - assert calls[0]["mode"] == "current_gpu_first" + assert not calls[0]["update_prefill_route_counter"] + assert calls[0]["mode"] == "global_first" def test_eplb_prefill_dispatch_consumes_physical_ids_and_event(monkeypatch): @@ -1152,7 +1178,7 @@ def dispatch(self, _qinput, **kwargs): impl.quant_method = object() runtime = _test_moe_impl( eplb=True, - route_counter=torch.zeros((128,), dtype=torch.int64), + prefill_route_counter=torch.zeros((24, 128), dtype=torch.int64), recording=True, ) _set_deepgemm_runtime(impl, runtime) @@ -1198,8 +1224,8 @@ def repair(**kwargs): assert topk_idx is physical_ids assert len(repair_calls) == 1 assert repair_calls[0]["logical_topk_ids"] is logical_ids - assert repair_calls[0]["update_logical_expert_counter"] - assert repair_calls[0]["mode"] == "current_gpu_first" + assert repair_calls[0]["update_prefill_route_counter"] + assert repair_calls[0]["mode"] == "global_first" assert calls[0]["topk_idx"] is physical_ids assert calls[0]["topk_idx"].dtype is torch.long assert calls[0]["previous_event"] is caller_event @@ -1264,7 +1290,8 @@ def cpu_zeros(*shape, **kwargs): assert impl.num_redundant_experts_per_rank == 1 assert impl.num_total_physical_experts == 6 - assert impl.route_counter.shape == (4,) + assert impl.prefill_route_counter.shape == (24, 4) + assert impl.prefill_route_sample_index.shape == (2,) assert impl.recording assert impl.local_logics_expert_ids_list == [0, 1, 2] assert not hasattr(impl, "initial_local_expert_ids_by_rank") @@ -1353,13 +1380,13 @@ def repair(**kwargs): return physical_ids monkeypatch.setattr(deepgemm_module, "eplb_repair_topk_ids", repair) - weights, selected = impl._prepare_expert_execution(torch.ones((1, 2)), logical_ids) + weights, selected = impl._prepare_expert_execution(torch.ones((1, 2)), logical_ids, is_prefill=True) assert weights.tolist() == [[1.0, 1.0]] assert selected is physical_ids assert calls[0]["logical_topk_ids"] is logical_ids - assert calls[0]["update_logical_expert_counter"] - assert calls[0]["mode"] == "current_gpu_first" + assert calls[0]["update_prefill_route_counter"] + assert calls[0]["mode"] == "global_first" def test_decode_masked_group_gemm_uses_all_physical_rows_when_eplb_is_enabled( @@ -1636,7 +1663,7 @@ def publish_layer_metadata(_layer_index): manager._commit_transfer = commit_transfer manager._publish_layer_metadata = publish_layer_metadata - manager._clear_route_counters = lambda: cleared_route_counters.append(True) + manager._clear_prefill_route_samples = lambda: cleared_route_counters.append(True) used_streams = [] overlap_stream = object() @@ -1723,7 +1750,7 @@ def test_manager_returns_to_collecting_after_reaching_rebalance_limit(): manager.pending_transfer_batches = [] manager.max_rebalance_count = 1 manager.completed_rebalance_count = 0 - manager._clear_route_counters = lambda: None + manager._clear_prefill_route_samples = lambda: None persisted_placements = [] manager._persist_current_placement = lambda: persisted_placements.append(manager.current_placement) @@ -1838,7 +1865,7 @@ def test_manager_step_advances_inflight_transfer(): def test_manager_evaluates_only_after_entering_evaluating_state(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - route_counter = torch.tensor([1, 2], dtype=torch.int64) + prefill_route_counter = torch.tensor([[1, 2]], dtype=torch.int64) local_token_counts = [] manager.state = manager_module.EPLBManagerState.COLLECTING manager.global_rank = 1 @@ -1846,7 +1873,7 @@ def test_manager_evaluates_only_after_entering_evaluating_state(monkeypatch): manager.step_interval = 3 manager.next_evaluation_step = 3 manager.num_logical_experts = 2 - manager._eplb_impls = [SimpleNamespace(route_counter=route_counter)] + manager._eplb_impls = [SimpleNamespace(prefill_route_counter=prefill_route_counter)] manager.world_size = 1 manager.control_group = object() manager.max_rebalance_count = -1 @@ -1872,7 +1899,7 @@ def all_gather_object(output, local_token_count, **_kwargs): assert manager.state is manager_module.EPLBManagerState.COLLECTING assert local_token_counts == [3] assert manager.next_evaluation_step == 6 - assert torch.equal(route_counter, torch.tensor([1, 2], dtype=torch.int64)) + assert torch.equal(prefill_route_counter, torch.tensor([[1, 2]], dtype=torch.int64)) def test_manager_step_uses_explicit_state_instead_of_pending_work(): @@ -1984,7 +2011,7 @@ def gather_finished(output, local_finished, **_kwargs): def test_manager_evaluation_with_insufficient_tokens_returns_to_collecting(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.EVALUATING - manager._eplb_impls = [SimpleNamespace(route_counter=torch.full((4,), 255, dtype=torch.int64))] + manager._eplb_impls = [SimpleNamespace(prefill_route_counter=torch.full((1, 4), 255, dtype=torch.int64))] manager.num_logical_experts = 4 manager.steps = 11 manager.step_interval = 20 @@ -2014,7 +2041,7 @@ def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): manager.state = manager_module.EPLBManagerState.EVALUATING manager.global_rank = 0 manager.num_logical_experts = 4 - manager._eplb_impls = [SimpleNamespace(route_counter=local_load[0])] + manager._eplb_impls = [SimpleNamespace(prefill_route_counter=local_load)] manager.world_size = 1 manager.control_group = object() manager.max_rebalance_count = -1 @@ -2075,11 +2102,11 @@ def test_manager_keeps_reporting_after_reaching_rebalance_limit(monkeypatch): manager.state = manager_module.EPLBManagerState.EVALUATING manager.global_rank = 0 manager.num_logical_experts = 4 - manager._eplb_impls = [SimpleNamespace(route_counter=local_load[0])] + manager._eplb_impls = [SimpleNamespace(prefill_route_counter=local_load)] manager.max_rebalance_count = 1 manager.completed_rebalance_count = 1 manager._publish_expert_load_metric = lambda load: published_loads.append(load) - manager._clear_route_counters = lambda: cleared_counters.append(True) + manager._clear_prefill_route_samples = lambda: cleared_counters.append(True) monkeypatch.setattr( manager_module.dist, "all_gather_object", @@ -2318,10 +2345,14 @@ def test_manager_rejects_sm100_before_initialization(monkeypatch): manager_module.EPLBManager(type("Model", (), {})()) -def test_manager_clears_all_route_counters_on_overlap_stream(monkeypatch): +def test_manager_clears_all_prefill_route_samples_on_overlap_stream(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - counters = [torch.tensor([1, 2]), torch.tensor([3, 4])] - manager._eplb_impls = [SimpleNamespace(route_counter=counter) for counter in counters] + counters = [torch.tensor([[1, 2]]), torch.tensor([[3, 4]])] + sample_indices = [torch.tensor([7, 3]), torch.tensor([9, 4])] + manager._eplb_impls = [ + SimpleNamespace(prefill_route_counter=counter, prefill_route_sample_index=sample_index) + for counter, sample_index in zip(counters, sample_indices) + ] overlap_stream = object() used_streams = [] monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) @@ -2331,10 +2362,11 @@ def test_manager_clears_all_route_counters_on_overlap_stream(monkeypatch): lambda stream: (used_streams.append(stream) or nullcontext()), ) - manager._clear_route_counters() + manager._clear_prefill_route_samples() assert used_streams == [overlap_stream] assert all(torch.count_nonzero(counter) == 0 for counter in counters) + assert all(torch.count_nonzero(sample_index) == 0 for sample_index in sample_indices) def test_manager_initializes_without_transfer_task(monkeypatch): @@ -2350,7 +2382,7 @@ def test_manager_initializes_without_transfer_task(monkeypatch): num_logical_experts=4, world_size=2, num_redundant_experts_per_rank=2, - route_counter=torch.zeros((4,), dtype=torch.int64), + prefill_route_counter=torch.zeros((24, 4), dtype=torch.int64), ), }, )() @@ -2365,7 +2397,7 @@ def test_manager_initializes_without_transfer_task(monkeypatch): clear_calls = [] monkeypatch.setattr( manager_module.EPLBManager, - "_clear_route_counters", + "_clear_prefill_route_samples", lambda manager: clear_calls.append(manager), ) monkeypatch.setattr(manager_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=1234)) @@ -2419,10 +2451,10 @@ def all_gather_object(output, local_expert_ids_by_layer, group): @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") -@pytest.mark.parametrize("update_logical_expert_counter", [False, True]) +@pytest.mark.parametrize("update_prefill_route_counter", [False, True]) @pytest.mark.parametrize("tokens", [1, 32]) @pytest.mark.parametrize("mode", ["current_gpu_first", "current_node_first", "global_first"]) -def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tokens, mode): +def test_eplb_repair_topk_ids_maps_and_counts(update_prefill_route_counter, tokens, mode): from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( eplb_repair_topk_ids, ) @@ -2466,7 +2498,8 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok ), dim=1, ) - counter = torch.zeros((experts,), dtype=torch.int64, device="cuda") + counter = torch.zeros((24, experts), dtype=torch.int64, device="cuda") + sample_index = torch.zeros((2,), dtype=torch.int64, device="cuda") expected_counter = torch.zeros_like(counter) logical_ids_long = logical_ids.to(torch.long) @@ -2490,8 +2523,8 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok hash_values ^= hash_values >> 16 replica_indices = hash_values % num_preferred_replicas[logical_ids_long].to(torch.int64) expected_ids = logical_to_physical[logical_ids_long, replica_indices + 3] - if update_logical_expert_counter: - expected_counter.scatter_add_( + if update_prefill_route_counter: + expected_counter[0].scatter_add_( 0, logical_ids.reshape(-1).to(torch.long), torch.ones(logical_ids.numel(), dtype=torch.int64, device="cuda"), @@ -2500,8 +2533,9 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok physical_ids = eplb_repair_topk_ids( logical_topk_ids=logical_ids, logical_to_physical_map=logical_to_physical, - logical_expert_counter=counter, - update_logical_expert_counter=update_logical_expert_counter, + prefill_route_counter=counter, + prefill_route_sample_index=sample_index, + update_prefill_route_counter=update_prefill_route_counter, mode=mode, ) torch.cuda.synchronize() @@ -2509,6 +2543,56 @@ def test_eplb_repair_topk_ids_maps_and_counts(update_logical_expert_counter, tok assert torch.equal(logical_ids, original_logical_ids) assert torch.equal(physical_ids, expected_ids) assert torch.equal(counter, expected_counter) + assert sample_index.tolist() == ([1, 0] if update_prefill_route_counter else [0, 0]) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +def test_eplb_prefill_route_counter_multigrid_ring_wrap_is_exact(): + from lightllm.common.basemodel.triton_kernel.fused_moe.eplb_topk_ids import ( + eplb_repair_topk_ids, + ) + + capacity = 24 + experts = 256 + tokens = 1025 + topk = 4 + num_samples = capacity + 2 + base_ids = torch.arange(tokens * topk, dtype=torch.int32, device="cuda").view(tokens, topk) + logical_experts = torch.arange(experts, dtype=torch.int32, device="cuda") + logical_to_physical = torch.stack( + ( + torch.ones_like(logical_experts), + torch.ones_like(logical_experts), + torch.ones_like(logical_experts), + logical_experts, + ), + dim=1, + ) + counter = torch.zeros((capacity, experts), dtype=torch.int64, device="cuda") + sample_index = torch.zeros((2,), dtype=torch.int64, device="cuda") + expected = torch.zeros_like(counter) + + for sample in range(num_samples): + logical_ids = (base_ids + sample) % experts + physical_ids = eplb_repair_topk_ids( + logical_topk_ids=logical_ids, + logical_to_physical_map=logical_to_physical, + prefill_route_counter=counter, + prefill_route_sample_index=sample_index, + update_prefill_route_counter=True, + mode="global_first", + ) + expected[sample % capacity] = torch.bincount( + logical_ids.reshape(-1).to(torch.long), + minlength=experts, + ) + assert torch.equal(physical_ids, logical_ids) + + torch.cuda.synchronize() + + assert sample_index.tolist() == [num_samples, 0] + assert torch.equal(counter, expected) + assert int(counter.sum().item()) == capacity * tokens * topk @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") @@ -2528,13 +2612,15 @@ def test_eplb_repair_topk_ids_spreads_strided_expert_tokens(): dtype=torch.int32, device="cuda", ) - counter = torch.zeros((2,), dtype=torch.int64, device="cuda") + counter = torch.zeros((24, 2), dtype=torch.int64, device="cuda") + sample_index = torch.zeros((2,), dtype=torch.int64, device="cuda") physical_ids = eplb_repair_topk_ids( logical_topk_ids=logical_ids, logical_to_physical_map=logical_to_physical, - logical_expert_counter=counter, - update_logical_expert_counter=False, + prefill_route_counter=counter, + prefill_route_sample_index=sample_index, + update_prefill_route_counter=False, mode="global_first", ) torch.cuda.synchronize() @@ -2562,13 +2648,15 @@ def test_eplb_replica_hash_is_uniform_across_tokens_and_experts(num_replicas): device="cuda", ) logical_to_physical = torch.cat((replica_counts, physical_ids), dim=1) - counter = torch.zeros((num_experts,), dtype=torch.int64, device="cuda") + counter = torch.zeros((24, num_experts), dtype=torch.int64, device="cuda") + sample_index = torch.zeros((2,), dtype=torch.int64, device="cuda") routed_physical_ids = eplb_repair_topk_ids( logical_topk_ids=logical_ids, logical_to_physical_map=logical_to_physical, - logical_expert_counter=counter, - update_logical_expert_counter=False, + prefill_route_counter=counter, + prefill_route_sample_index=sample_index, + update_prefill_route_counter=False, mode="global_first", ) torch.cuda.synchronize() @@ -2604,7 +2692,8 @@ def test_eplb_repair_topk_ids_empty_input_skips_kernel(): experts = 64 logical_ids = torch.empty((0, 4), dtype=torch.int32, device="cuda") - counter = torch.zeros((experts,), dtype=torch.int64, device="cuda") + counter = torch.zeros((24, experts), dtype=torch.int64, device="cuda") + sample_index = torch.zeros((2,), dtype=torch.int64, device="cuda") logical_to_physical = torch.stack( ( torch.ones((experts,), dtype=torch.int32, device="cuda"), @@ -2617,14 +2706,16 @@ def test_eplb_repair_topk_ids_empty_input_skips_kernel(): physical_ids = eplb_repair_topk_ids( logical_topk_ids=logical_ids, logical_to_physical_map=logical_to_physical, - logical_expert_counter=counter, - update_logical_expert_counter=True, + prefill_route_counter=counter, + prefill_route_sample_index=sample_index, + update_prefill_route_counter=True, mode="current_gpu_first", ) assert physical_ids.shape == (0, 4) assert physical_ids.dtype is torch.int32 assert torch.equal(counter, torch.zeros_like(counter)) + assert torch.equal(sample_index, torch.zeros_like(sample_index)) @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") @@ -2635,13 +2726,15 @@ def test_eplb_repair_topk_ids_rejects_unknown_dispatch_mode(): logical_ids = torch.empty((0, 1), dtype=torch.int32, device="cuda") logical_to_physical = torch.tensor([[1, 1, 1, 0]], dtype=torch.int32, device="cuda") - counter = torch.zeros((1,), dtype=torch.int64, device="cuda") + counter = torch.zeros((24, 1), dtype=torch.int64, device="cuda") + sample_index = torch.zeros((2,), dtype=torch.int64, device="cuda") with pytest.raises(AssertionError, match="unsupported EPLB dispatch mode"): eplb_repair_topk_ids( logical_topk_ids=logical_ids, logical_to_physical_map=logical_to_physical, - logical_expert_counter=counter, - update_logical_expert_counter=False, + prefill_route_counter=counter, + prefill_route_sample_index=sample_index, + update_prefill_route_counter=False, mode="unknown", ) From b3beb21f863f1d063b2c0891629d41d481ec760c Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sat, 26 Sep 2026 11:10:05 +0000 Subject: [PATCH 68/72] feat(eplb): report placement imbalance metrics --- docs/CN/source/framework/eplb.md | 33 +++- lightllm/server/metrics/metrics.py | 25 ++- .../model_infer/mode_backend/eplb/metrics.py | 142 ++++++++++++++++ .../mode_backend/eplb/runtime_manager.py | 48 +++--- unit_tests/common/fused_moe/test_eplb.py | 160 +++++++++++++++--- 5 files changed, 354 insertions(+), 54 deletions(-) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/metrics.py diff --git a/docs/CN/source/framework/eplb.md b/docs/CN/source/framework/eplb.md index 052726a8c5..3161f1e0a1 100644 --- a/docs/CN/source/framework/eplb.md +++ b/docs/CN/source/framework/eplb.md @@ -247,7 +247,7 @@ ready 发布使用 `release`、等待方使用 `acquire`,最后完成信号使 ### 5.7 manager 聚合与重置 -manager 在安全推理边界把各层 `[24, E]` 环形缓冲区沿第 0 维求和并复制到 CPU,得到 planner 使用的 `[layer, logical_expert]` 负载。如果样本量不足或规划结果未改变布局,不主动清空缓冲区;后续 prefill 会继续写入,并在容量用满后滚动覆盖最旧行。 +manager 在安全推理边界一次性堆叠各层 `[24, E]` 环形缓冲区,在 GPU 上沿 sample 维求和后复制到 CPU,得到 planner 使用的 `[layer, logical_expert]` 负载。如果样本量不足或规划结果未改变布局,不主动清空缓冲区;后续 prefill 会继续写入,并在容量用满后滚动覆盖最旧行。 初始化、成功切换到新布局,以及达到重排次数上限后开始下一轮指标窗口时,manager 会同时清零 `prefill_route_counter` 和 `prefill_route_sample_index`。清零提交到 overlap stream,自然排在此前 forward 之后、后续 forward 之前,不需要额外的全设备同步。 @@ -437,13 +437,38 @@ DP 随机分流下每个 rank 的路由统计都是对全局路由分布的无 ## 11. 指标与运行行为 -rank 0 周期性上报: +rank 0 周期性上报 logical expert 路由分布: ```text -lightllm_eplb_topk_expert_imbalance_ratio +lightllm_eplb_topk_expert_imbalance_ratio_p25 +lightllm_eplb_topk_expert_imbalance_ratio_p50 +lightllm_eplb_topk_expert_imbalance_ratio_p100 ``` -该指标先计算每层 `max(expert_load) / mean(expert_load)`,再对有效层求平均。值越接近 1,表示观测窗口内的逻辑专家负载越均衡。 +manager 先对每个有效层计算 `max(expert_load) / mean(expert_load)`,过滤没有采样负载的层,再对所有层的比值排序并使用 nearest-rank 位置取值: + +- P25:较均衡的四分之一位置,可观察大多数浅层或稳定层的基线; +- P50:中位层,用于描述典型 MoE 层的路由倾斜程度; +- P100:最大值,即当前窗口中最不均衡的层。 + +这三个值都以 `1` 表示完全均衡。例如 P50 为 `1.8`,表示中位层最热 logical expert 的 token 数是该层专家平均值的 1.8 倍。 + +当全局样本量达到规划阈值且 planner 产生目标布局后,rank 0 使用 planner 实际消费的同一份 `global_load` 上报: + +```text +lightllm_prefill_ep_compute_critical_overhead_ratio_before_rebalance +lightllm_prefill_ep_compute_critical_overhead_ratio_after_rebalance +``` + +这两个指标参考 `eplb2` 的关键路径计算开销定义,但不维护独立的 compute counter、后台 monitor 线程和额外通信组。manager 假设同一 logical expert 的流量由 hash 均匀分配给全部 physical 副本,并按 128 token 对每个副本的估算负载向上对齐。`before_rebalance` 使用当前布局,`after_rebalance` 使用 planner 给出的目标布局;二者的输入负载完全相同,可以直接衡量预计的重排收益。如果 planner 判断布局无需改变,两个值应相同。 + +每层先计算最繁忙 rank 相对平均 rank 的额外负载,最后跨层汇总: + +```text +overhead_ratio = sum(max_rank_load - mean_rank_load) / sum(mean_rank_load) +``` + +因此不同层的热点 rank 不会互相抵消。指标为 `0` 表示估算的 EP rank 负载完全均衡,`0.3` 表示最慢 rank 造成的关键路径计算量比理想均衡状态高约 30%。它是基于聚合 logical route 和均匀副本分发假设的布局质量估算值,不是 DeepEP 接收缓冲区的实测 compute load,也不表示每次 prefill 的瞬时开销。只有实际执行 placement 规划时这两个 gauge 才会更新;其余时间保留最近一次规划结果。 `--eplb_rebalance_count` 的行为如下: diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index 0d53669a16..a0db50c92c 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -32,9 +32,22 @@ "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", "lightllm_gen_throughput": "Generation throughput of latest completed request (tokens/s)", "lightllm_num_running_reqs": "Number of running requests", - "lightllm_eplb_topk_expert_imbalance_ratio": ( - "Maximum routed token count divided by the mean across logical experts, averaged across MoE layers in the " - "accumulated EPLB routing sample" + "lightllm_prefill_ep_compute_critical_overhead_ratio_before_rebalance": ( + "Estimated excess critical EP-rank compute divided by balanced compute for the global load used by the latest " + "EPLB placement plan, evaluated before rebalance; 0.3 means 30% overhead" + ), + "lightllm_prefill_ep_compute_critical_overhead_ratio_after_rebalance": ( + "Estimated excess critical EP-rank compute divided by balanced compute for the global load used by the latest " + "EPLB placement plan, evaluated on the planned placement; 0.3 means 30% overhead" + ), + "lightllm_eplb_topk_expert_imbalance_ratio_p25": ( + "P25 across MoE layers of maximum-to-mean logical expert routed-token load" + ), + "lightllm_eplb_topk_expert_imbalance_ratio_p50": ( + "P50 across MoE layers of maximum-to-mean logical expert routed-token load" + ), + "lightllm_eplb_topk_expert_imbalance_ratio_p100": ( + "P100 across MoE layers of maximum-to-mean logical expert routed-token load" ), } @@ -115,7 +128,11 @@ def init_metrics(self, args): self.create_gauge("lightllm_cache_hit_rate") self.create_gauge("lightllm_gen_throughput") self.create_gauge("lightllm_num_running_reqs") - self.create_gauge("lightllm_eplb_topk_expert_imbalance_ratio") + self.create_gauge("lightllm_prefill_ep_compute_critical_overhead_ratio_before_rebalance") + self.create_gauge("lightllm_prefill_ep_compute_critical_overhead_ratio_after_rebalance") + self.create_gauge("lightllm_eplb_topk_expert_imbalance_ratio_p25") + self.create_gauge("lightllm_eplb_topk_expert_imbalance_ratio_p50") + self.create_gauge("lightllm_eplb_topk_expert_imbalance_ratio_p100") def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py b/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py new file mode 100644 index 0000000000..2323c3273c --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py @@ -0,0 +1,142 @@ +"""计算并上报 EPLB 的 logical expert 与 EP rank 负载指标。""" + +import numpy as np +import torch + +from lightllm.server.metrics.manager import MetricClient + +from .placement import ExpertPlacement + + +COMPUTE_CRITICAL_OVERHEAD_RATIO_BEFORE_REBALANCE_METRIC = ( + "lightllm_prefill_ep_compute_critical_overhead_ratio_before_rebalance" +) +COMPUTE_CRITICAL_OVERHEAD_RATIO_AFTER_REBALANCE_METRIC = ( + "lightllm_prefill_ep_compute_critical_overhead_ratio_after_rebalance" +) +EXPERT_IMBALANCE_RATIO_METRICS = { + 25: "lightllm_eplb_topk_expert_imbalance_ratio_p25", + 50: "lightllm_eplb_topk_expert_imbalance_ratio_p50", + 100: "lightllm_eplb_topk_expert_imbalance_ratio_p100", +} + + +def logical_expert_imbalance_percentiles(*, expert_load: torch.Tensor) -> dict[int, float]: + """统计各层 ``最热 logical expert / 本层平均负载`` 的分位数。""" + assert expert_load.ndim == 2 and expert_load.numel() > 0 + + expert_load = expert_load.to(torch.float64) + mean_load_by_layer = expert_load.mean(dim=1) + + # 没有 token 的层不参与分位数计算,避免产生 0 / 0。 + has_route_load = mean_load_by_layer > 0 + if not torch.any(has_route_load): + return {percentile: 0.0 for percentile in EXPERT_IMBALANCE_RATIO_METRICS} + + hottest_expert_load = expert_load.max(dim=1).values + imbalance_ratio_by_layer = hottest_expert_load[has_route_load] / mean_load_by_layer[has_route_load] + percentiles = tuple(EXPERT_IMBALANCE_RATIO_METRICS) + + # inverted_cdf 就是 nearest-rank 定义;默认 linear 或 nearest 插值都会改变 + # 层数较少时的 P25/P50 语义。 + percentile_values = np.percentile( + a=imbalance_ratio_by_layer.numpy(), + q=percentiles, + method="inverted_cdf", + ) + return dict(zip(percentiles, percentile_values.tolist())) + + +def publish_expert_load_metrics(*, metric_client: MetricClient, expert_load: torch.Tensor) -> None: + """上报 logical expert 层间不均衡分位数。""" + imbalance_percentiles = logical_expert_imbalance_percentiles(expert_load=expert_load) + for percentile, metric_name in EXPERT_IMBALANCE_RATIO_METRICS.items(): + metric_client.gauge_set( + name=metric_name, + value=imbalance_percentiles[percentile], + ) + + +def compute_critical_overhead_ratio( + *, + logical_expert_load: torch.Tensor, + placement: ExpertPlacement, + expert_alignment: int, +) -> float: + """估算 placement 相对理想 rank 均衡状态的关键路径额外计算比例。 + + 假设同一 logical expert 的流量由 hash 均匀分配给所有 physical 副本, + 并按 DeepEP 的 expert alignment 对每个副本负载向上取整。每层单独选择 + 最繁忙 rank,避免不同层的热点 rank 在汇总时相互抵消。 + """ + assert logical_expert_load.ndim == 2 and logical_expert_load.numel() > 0 + assert expert_alignment > 0 + + num_layers, num_logical_experts = logical_expert_load.shape + + # placement: [layer, rank, local physical expert] + placement_tensor = torch.tensor(placement, dtype=torch.int64) + assert placement_tensor.ndim == 3 and placement_tensor.shape[0] == num_layers + + # 将每个 physical slot 中保存的 logical expert ID 转成 one-hot: + # [layer, rank, physical slot] -> [layer, rank, physical slot, logical expert]。 + # 例如 slot 中保存 expert 2,就会在 logical expert 维得到 [0, 0, 1, ...]。 + expert_mask_by_physical_slot = torch.nn.functional.one_hot( + placement_tensor, + num_classes=num_logical_experts, + ) + + # 沿 physical slot 维求和,得到每个 rank 持有的专家副本数:[layer, rank, logical expert]。 + replicas_per_rank = expert_mask_by_physical_slot.sum(dim=2).to(torch.float64) + + # 再沿 rank 维求和,得到每个 logical expert 的全局副本数:[layer, logical expert]。 + replica_count_per_expert = replicas_per_rank.sum(dim=1) + assert torch.all(replica_count_per_expert > 0) + + # logical_expert_load: [layer, logical expert] + # hash 均匀分流后,每个 physical 副本承担 logical expert 总负载的 1/N。 + load_per_replica = logical_expert_load.to(torch.float64) / replica_count_per_expert + aligned_load_per_replica = torch.ceil(load_per_replica / expert_alignment) * expert_alignment + + # 广播为 [layer, rank, logical expert] 后沿 expert 维求和,得到每层各 rank + # 的估算计算量:[layer, rank]。 + estimated_rank_load = (replicas_per_rank * aligned_load_per_replica.unsqueeze(dim=1)).sum(dim=2) + + mean_rank_load_by_layer = estimated_rank_load.mean(dim=1) + critical_rank_load_by_layer = estimated_rank_load.max(dim=1).values + total_balanced_compute = mean_rank_load_by_layer.sum() + if total_balanced_compute == 0: + return 0.0 + + total_critical_overhead = (critical_rank_load_by_layer - mean_rank_load_by_layer).sum() + return float((total_critical_overhead / total_balanced_compute).item()) + + +def publish_rebalance_compute_metrics( + *, + metric_client: MetricClient, + global_load: torch.Tensor, + current_placement: ExpertPlacement, + target_placement: ExpertPlacement, + expert_alignment: int, +) -> None: + """使用同一份 planner 输入上报重排前后的关键路径开销。""" + before_rebalance_ratio = compute_critical_overhead_ratio( + logical_expert_load=global_load, + placement=current_placement, + expert_alignment=expert_alignment, + ) + after_rebalance_ratio = compute_critical_overhead_ratio( + logical_expert_load=global_load, + placement=target_placement, + expert_alignment=expert_alignment, + ) + + metric_client.gauge_set( + name=COMPUTE_CRITICAL_OVERHEAD_RATIO_BEFORE_REBALANCE_METRIC, + value=before_rebalance_ratio, + ) + metric_client.gauge_set( + name=COMPUTE_CRITICAL_OVERHEAD_RATIO_AFTER_REBALANCE_METRIC, + value=after_rebalance_ratio, + ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index bf3953d331..48d185321d 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -20,6 +20,7 @@ from lightllm.utils.log_utils import init_logger from lightllm.utils.shm_port_args import get_shm_port_args +from . import metrics as eplb_metrics from .async_transfer_planner import EPLBTransferPlanner from .expert_transfer import ( EPLBTransferInfo, @@ -37,7 +38,6 @@ logger = init_logger(__name__) EPLB_EXPERT_ALIGNMENT = 128 EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT = 128 -EPLB_EXPERT_IMBALANCE_RATIO_METRIC = "lightllm_eplb_topk_expert_imbalance_ratio" class EPLBManagerState(Enum): @@ -236,14 +236,18 @@ def _step_evaluating(self) -> None: if len({counter.shape[0] for counter in counters}) != 1: raise RuntimeError("EPLB prefill route counter capacities must match across layers") - # 将各层环形样本聚合并复制到 CPU,后续规划统一使用这份快照。 + # 在 GPU 上沿 sample 维聚合各层的环形样本,再一次性复制到 CPU, + # 得到 planner 使用的 [layer, logical_expert] 负载,避免逐层发起 + # GPU -> CPU 拷贝。 # 此处先不清零 GPU 样本:如果样本不足或无需迁移,下一周期会继续 # 滚动覆盖最旧行;达到重排上限或成功切换到新布局后才重置窗口。本轮 # 异步规划使用独立的 CPU 快照,不会与后续的 atomic add 竞争。 - # 当前 planner 仍消费 [layer, logical_expert] 聚合负载;先对 24 行采样 - # 求和保持现有规划语义。后续逐样本规划可以直接在这里保留 sample 维。 - local_load = torch.stack([counter.detach().sum(dim=0).cpu() for counter in counters]) - self._publish_expert_load_metric(local_load) + local_load = torch.stack(counters).sum(dim=1).detach().cpu() + if self.global_rank == 0: + eplb_metrics.publish_expert_load_metrics( + metric_client=self.metric_client, + expert_load=local_load, + ) # 达到重排次数上限后仍保留周期性负载上报,但不再执行后续的跨 rank # 通信和布局规划。清空本轮样本,使下一次指标对应新的采样窗口。 @@ -292,6 +296,9 @@ def _step_plan_placement(self) -> None: self.state = EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED if self.global_rank == 0: + # 保留 planner 实际消费的全局负载快照。目标布局产生后,使用同一份 + # 输入分别评估 current/target placement,确保 before/after 可直接比较。 + self._planning_global_load = global_load self._plan_task = EPLBPlanTask( self.planner, global_load, @@ -313,6 +320,14 @@ def _step_wait_plan_placement_finished(self) -> None: return if self.global_rank == 0: + eplb_metrics.publish_rebalance_compute_metrics( + metric_client=self.metric_client, + global_load=self._planning_global_load, + current_placement=self.current_placement, + target_placement=placement, + expert_alignment=EPLB_EXPERT_ALIGNMENT, + ) + del self._planning_global_load del self._plan_task if placement == self.current_placement: @@ -487,14 +502,6 @@ def _publish_layer_metadata(self, layer_index: int) -> None: ) layer_impl.logical_to_physical_map.copy_(logical_to_physical_map, non_blocking=True) - def _publish_expert_load_metric(self, local_load: torch.Tensor) -> None: - if self.global_rank != 0: - return - self.metric_client.gauge_set( - EPLB_EXPERT_IMBALANCE_RATIO_METRIC, - _expert_load_imbalance_ratio(local_load), - ) - def _clear_prefill_route_samples(self) -> None: """在 overlap stream 上清空所有层的 prefill 路由样本和设备端索引。""" from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -528,16 +535,3 @@ def _find_fused_moe_weights(model: TpPartBaseModel) -> List[FusedMoeWeight]: if isinstance(weight, FusedMoeWeight) and weight.enable_ep_moe: weights.append(weight) return weights - - -def _expert_load_imbalance_ratio(expert_load: torch.Tensor) -> float: - """计算各层逻辑专家最大 token 数与平均值之比,再对所有层取平均。""" - if expert_load.ndim != 2: - raise ValueError("expert_load must be [layers, logical_experts]") - expert_load = expert_load.to(torch.float64) - layer_means = expert_load.mean(dim=1) - valid_layers = layer_means > 0 - if not torch.any(valid_layers): - return 0.0 - ratios = expert_load.max(dim=1).values[valid_layers] / layer_means[valid_layers] - return float(ratios.mean().item()) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 596e4fdce8..d47187485a 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -19,6 +19,7 @@ from lightllm.server.router.model_infer.mode_backend.eplb import ( runtime_manager as manager_module, ) +from lightllm.server.router.model_infer.mode_backend.eplb import metrics as eplb_metrics from lightllm.server.router.model_infer.mode_backend.eplb import ( placement_plan_task as plan_module, ) @@ -907,40 +908,118 @@ def fail(*_args): assert logs == ["EPLB transfer planning failed"] -def test_expert_load_imbalance_ratio_averages_layer_ratios(): - global_load = torch.tensor( +def test_compute_critical_overhead_ratio_estimates_rank_pressure(): + load = torch.tensor([[384, 128, 128, 128]], dtype=torch.int64) + placement = [[[0, 1], [2, 3]]] + + ratio = eplb_metrics.compute_critical_overhead_ratio( + logical_expert_load=load, + placement=placement, + expert_alignment=128, + ) + + # rank loads are [512, 256], so excess critical / balanced is 128 / 384. + assert ratio == pytest.approx(1 / 3) + + +def test_compute_critical_overhead_ratio_preserves_layer_boundaries(): + load = torch.tensor( + [ + [1, 0], + [0, 1], + ], + dtype=torch.int64, + ) + + ratio = eplb_metrics.compute_critical_overhead_ratio( + logical_expert_load=load, + placement=[ + [[0], [1]], + [[0], [1]], + ], + expert_alignment=128, + ) + + # 两层的热点 rank 相反;逐层取关键路径时仍各有 100% 开销,不能相互抵消。 + assert ratio == pytest.approx(1.0) + + +def test_compute_critical_overhead_ratio_is_zero_without_load(): + assert ( + eplb_metrics.compute_critical_overhead_ratio( + logical_expert_load=torch.zeros((1, 2), dtype=torch.int64), + placement=[[[0], [1]]], + expert_alignment=128, + ) + == 0.0 + ) + + +def test_logical_expert_imbalance_percentiles_report_layer_distribution(): + load = torch.tensor( [ - [2, 4, 6], - [10, 10, 10], + [3, 3, 3, 3], + [6, 2, 2, 2], + [9, 1, 1, 1], + [12, 0, 0, 0], + [0, 0, 0, 0], ], dtype=torch.int64, ) - ratio = manager_module._expert_load_imbalance_ratio(global_load) + assert eplb_metrics.logical_expert_imbalance_percentiles(expert_load=load) == { + 25: pytest.approx(1.0), + 50: pytest.approx(2.0), + 100: pytest.approx(4.0), + } + - assert ratio == pytest.approx(1.25) +def test_logical_expert_imbalance_percentiles_are_zero_without_load(): + assert eplb_metrics.logical_expert_imbalance_percentiles(expert_load=torch.zeros((3, 4), dtype=torch.int64)) == { + 25: 0.0, + 50: 0.0, + 100: 0.0, + } -def test_manager_publishes_expert_load_metrics_from_rank_zero(): +def test_publish_expert_load_metrics(): calls = [] - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.global_rank = 0 - manager.metric_client = SimpleNamespace(gauge_set=lambda name, value: calls.append((name, value))) + metric_client = SimpleNamespace(gauge_set=lambda name, value: calls.append((name, value))) - manager._publish_expert_load_metric(torch.tensor([[2, 4, 6], [10, 10, 10]])) + eplb_metrics.publish_expert_load_metrics( + metric_client=metric_client, + expert_load=torch.tensor([[384, 128, 128, 128]]), + ) assert calls == [ - (manager_module.EPLB_EXPERT_IMBALANCE_RATIO_METRIC, 1.25), + (eplb_metrics.EXPERT_IMBALANCE_RATIO_METRICS[25], pytest.approx(2.0)), + (eplb_metrics.EXPERT_IMBALANCE_RATIO_METRICS[50], pytest.approx(2.0)), + (eplb_metrics.EXPERT_IMBALANCE_RATIO_METRICS[100], pytest.approx(2.0)), ] -def test_manager_does_not_publish_expert_load_metrics_from_other_ranks(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.global_rank = 1 +def test_publish_rebalance_compute_metrics_from_global_load(): + calls = [] + metric_client = SimpleNamespace(gauge_set=lambda name, value: calls.append((name, value))) - manager._publish_expert_load_metric(torch.tensor([[1, 2]])) + eplb_metrics.publish_rebalance_compute_metrics( + metric_client=metric_client, + global_load=torch.tensor([[384, 128, 128, 128]]), + current_placement=[[[0, 1, 2], [1, 2, 3]]], + target_placement=[[[0, 1, 2], [0, 2, 3]]], + expert_alignment=128, + ) - assert not hasattr(manager, "metric_client") + assert calls == [ + ( + eplb_metrics.COMPUTE_CRITICAL_OVERHEAD_RATIO_BEFORE_REBALANCE_METRIC, + pytest.approx(0.25), + ), + ( + eplb_metrics.COMPUTE_CRITICAL_OVERHEAD_RATIO_AFTER_REBALANCE_METRIC, + pytest.approx(0.0), + ), + ] def test_eplb_prefill_route_counter_has_24_samples_per_logical_expert(monkeypatch): @@ -1783,6 +1862,38 @@ def broadcast(values, **_kwargs): assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED +def test_wait_plan_finish_publishes_before_and_after_metrics(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.state = manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED + manager.global_rank = 0 + manager.control_group = object() + manager.current_placement = [[[0, 1], [2, 3]]] + target_placement = [[[0, 2], [1, 3]]] + manager._planning_global_load = torch.tensor([[4, 3, 2, 1]], dtype=torch.int64) + manager._plan_task = SimpleNamespace( + is_finished=lambda: True, + result=target_placement, + ) + manager.metric_client = object() + published = [] + monkeypatch.setattr( + eplb_metrics, + "publish_rebalance_compute_metrics", + lambda **kwargs: published.append((kwargs["global_load"], kwargs["target_placement"])), + ) + monkeypatch.setattr(manager_module.dist, "broadcast_object_list", lambda _values, **_kwargs: None) + + manager._step_wait_plan_placement_finished() + + assert manager.state is manager_module.EPLBManagerState.PLAN_TRANSFER + assert manager.target_placement is target_placement + assert len(published) == 1 + assert torch.equal(published[0][0], torch.tensor([[4, 3, 2, 1]], dtype=torch.int64)) + assert published[0][1] is target_placement + assert not hasattr(manager, "_planning_global_load") + assert not hasattr(manager, "_plan_task") + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_manager_transfer_task_commit_orders_live_weights_between_overlap_forwards( monkeypatch, @@ -2070,7 +2181,12 @@ def start(self): manager.planner = object() manager.current_placement = [[[0, 1, 2, 3]]] - manager._publish_expert_load_metric = lambda load: published_loads.append(load) + manager.metric_client = object() + monkeypatch.setattr( + eplb_metrics, + "publish_expert_load_metrics", + lambda **kwargs: published_loads.append(kwargs["expert_load"]), + ) monkeypatch.setattr(manager_module, "EPLBPlanTask", PlanTask) manager.step() @@ -2090,6 +2206,7 @@ def start(self): assert plan_tasks[0].current_placement is manager.current_placement assert plan_tasks[0].started assert manager._plan_task is plan_tasks[0] + assert torch.equal(manager._planning_global_load, local_load) assert not hasattr(manager, "_local_load") assert len(published_loads) == 1 @@ -2105,7 +2222,12 @@ def test_manager_keeps_reporting_after_reaching_rebalance_limit(monkeypatch): manager._eplb_impls = [SimpleNamespace(prefill_route_counter=local_load)] manager.max_rebalance_count = 1 manager.completed_rebalance_count = 1 - manager._publish_expert_load_metric = lambda load: published_loads.append(load) + manager.metric_client = object() + monkeypatch.setattr( + eplb_metrics, + "publish_expert_load_metrics", + lambda **kwargs: published_loads.append(kwargs["expert_load"]), + ) manager._clear_prefill_route_samples = lambda: cleared_counters.append(True) monkeypatch.setattr( manager_module.dist, From b94fd18358529a7b39206d93bf96f6ce4e8f50c2 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sat, 26 Sep 2026 12:39:00 +0000 Subject: [PATCH 69/72] refactor(eplb): simplify async task statuses --- .../eplb/async_transfer_planner.py | 19 ++++----------- .../mode_backend/eplb/expert_transfer.py | 23 ++++++------------- .../mode_backend/eplb/placement_plan_task.py | 19 ++++----------- unit_tests/common/fused_moe/test_eplb.py | 21 ++++++++--------- 4 files changed, 27 insertions(+), 55 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py index 3e3746a63c..e5dafaa560 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py @@ -2,7 +2,6 @@ import os import threading -from enum import Enum from typing import List, Optional from lightllm.utils.log_utils import init_logger @@ -13,14 +12,6 @@ logger = init_logger(__name__) -class TransferPlanStatus(Enum): - """异步传输规划器的生命周期状态。""" - - IDLE = "idle" - RUNNING = "running" - SUCCEEDED = "succeeded" - - class EPLBTransferPlanner: """在后台线程中逐层生成并按 layer 顺序拼接专家传输批次。 @@ -40,7 +31,7 @@ def __init__( self.target_placement = target_placement self.num_logical_experts = num_logical_experts self.world_size = world_size - self.status = TransferPlanStatus.IDLE + self.status = "idle" self.result: Optional[List[List[EPLBTransferInfo]]] = None self._thread = threading.Thread( target=self._run, @@ -50,13 +41,13 @@ def __init__( def start(self) -> None: """启动异步传输规划。""" - assert self.status is TransferPlanStatus.IDLE, "EPLB transfer planner has already been started" - self.status = TransferPlanStatus.RUNNING + assert self.status == "idle", "EPLB transfer planner has already been started" + self.status = "running" self._thread.start() def is_finished(self) -> bool: """返回全部层的传输批次是否已经生成。""" - return self.status is TransferPlanStatus.SUCCEEDED + return self.status == "succeeded" def _run(self) -> None: try: @@ -72,7 +63,7 @@ def _run(self) -> None: ) transfer_batches.extend(layer_transfer_batches) self.result = transfer_batches - self.status = TransferPlanStatus.SUCCEEDED + self.status = "succeeded" except BaseException: logger.exception("EPLB transfer planning failed") os._exit(1) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py index 67e0c8e901..2943b4fad6 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py @@ -4,7 +4,6 @@ import threading import zlib from dataclasses import dataclass -from enum import Enum from typing import List, Sequence import torch @@ -44,14 +43,6 @@ class EPLBTransferInfo: dest_local_expert_index: int -class TransferStatus(Enum): - """异步传输线程的生命周期状态。""" - - IDLE = "idle" - RUNNING = "running" - SUCCEEDED = "succeeded" - - class PinnedMemoryEPLBTransfer: """在后台线程中传输一个逻辑专家的全部权重张量。 @@ -63,8 +54,8 @@ class PinnedMemoryEPLBTransfer: 网络通信。其他 rank 不分配 pinned row,也不参与该专家的数据传输。 本类只负责异步传输,不修改 live 权重,也不更新路由 metadata。传输成功后, - ``status`` 会变为 :attr:`TransferStatus.SUCCEEDED`,收到的数据保存在 - ``tensor_buffers``。EPLBManager 在主循环的安全边界同步提交这些数据。 + ``status`` 会变为 ``"succeeded"``,收到的数据保存在 ``tensor_buffers``。 + EPLBManager 在主循环的安全边界同步提交这些数据。 每个对象只表示构造函数中 ``transfer_info`` 指定的一次传输。源 rank 直接 读取 ``source_local_expert_index`` 指定的物理行;目标物理槽位不属于传输 @@ -111,7 +102,7 @@ def __init__( ) self._device_to_host_stream: torch.cuda.Stream = torch.cuda.Stream(device=self._device) - self.status: TransferStatus = TransferStatus.IDLE + self.status = "idle" self._transfer_thread: threading.Thread = threading.Thread( target=self._run_transfer, name=f"eplb-transfer-layer-{transfer_info.layer_index}-expert-{transfer_info.source_logical_expert_id}", @@ -120,13 +111,13 @@ def __init__( def start(self) -> None: """启动构造函数中 transfer_info 描述的异步传输。""" - assert self.status is TransferStatus.IDLE, "EPLB transfer has already been started" - self.status = TransferStatus.RUNNING + assert self.status == "idle", "EPLB transfer has already been started" + self.status = "running" self._transfer_thread.start() def is_finished(self) -> bool: """返回后台传输是否已经成功完成。""" - return self.status is TransferStatus.SUCCEEDED + return self.status == "succeeded" def _run_transfer(self) -> None: """把指定专家的全部权重行传输到各 rank 的 pinned memory。""" @@ -162,7 +153,7 @@ def _run_transfer(self) -> None: group=self._p2p_group, tag=message_tag, ) - self.status = TransferStatus.SUCCEEDED + self.status = "succeeded" except BaseException: logger.exception("EPLB transfer failed") os._exit(1) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py index 8fc80d8574..b7b06bab27 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py @@ -2,7 +2,6 @@ import os import threading -from enum import Enum from typing import Optional import torch @@ -14,14 +13,6 @@ logger = init_logger(__name__) -class PlanTaskStatus(Enum): - """异步规划任务的生命周期状态。""" - - IDLE = "idle" - RUNNING = "running" - SUCCEEDED = "succeeded" - - class EPLBPlanTask: """在后台线程中根据全局专家负载生成新布局。""" @@ -34,7 +25,7 @@ def __init__( self.planner = planner self.global_load = global_load self.current_placement = current_placement - self.status = PlanTaskStatus.IDLE + self.status = "idle" self.result: Optional[ExpertPlacement] = None self._thread = threading.Thread( target=self._run, @@ -44,13 +35,13 @@ def __init__( def start(self) -> None: """启动异步规划。""" - assert self.status is PlanTaskStatus.IDLE, "EPLB plan task has already been started" - self.status = PlanTaskStatus.RUNNING + assert self.status == "idle", "EPLB plan task has already been started" + self.status = "running" self._thread.start() def is_finished(self) -> bool: """返回规划任务是否已经成功完成。""" - return self.status is PlanTaskStatus.SUCCEEDED + return self.status == "succeeded" def _run(self) -> None: try: @@ -58,7 +49,7 @@ def _run(self) -> None: self.global_load.tolist(), self.current_placement, ) - self.status = PlanTaskStatus.SUCCEEDED + self.status = "succeeded" except BaseException: logger.exception("EPLB planning failed") os._exit(1) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index d47187485a..9c12a37026 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -48,7 +48,6 @@ EPLBTransferInfo, ExpertTensorBuffer, PinnedMemoryEPLBTransfer, - TransferStatus, build_transfer_plan, ) @@ -828,7 +827,7 @@ def test_manager_delegates_distribution_planning_to_planner_class(): task._run() - assert task.status is plan_module.PlanTaskStatus.SUCCEEDED + assert task.status == "succeeded" assert task.result == planned_placement assert len(calls) == 1 assert calls[0][0] == logical_load.tolist() @@ -878,7 +877,7 @@ def build_plan(*args): planner._run() - assert planner.status is transfer_planner_module.TransferPlanStatus.SUCCEEDED + assert planner.status == "succeeded" assert planner.result == [[transfer_infos[0]], [transfer_infos[1]]] assert calls == [ (current_placement[0], target_placement[0], 0, 4, 2), @@ -1902,7 +1901,7 @@ class Transfer: def __init__(self, live, received, transfer_info): self.tensor_buffers = [ExpertTensorBuffer("weight", live, received)] self.transfer_info = transfer_info - self.status = TransferStatus.SUCCEEDED + self.status = "succeeded" def is_finished(self): return True @@ -2318,7 +2317,7 @@ def synchronize(self): torch.empty(2), ) ] - transfer.status = TransferStatus.RUNNING + transfer.status = "running" sends = [] monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: nullcontext()) @@ -2330,7 +2329,7 @@ def synchronize(self): transfer._run_transfer() - assert transfer.status is TransferStatus.SUCCEEDED + assert transfer.status == "succeeded" assert len(sends) == 1 assert torch.equal(sends[0][0], torch.tensor([3.0, 4.0])) expected_tag = transfer._build_p2p_message_tag("weight") @@ -2358,7 +2357,7 @@ def synchronize(self): torch.empty(2), ) ] - transfer.status = TransferStatus.RUNNING + transfer.status = "running" p2p_calls = [] monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: nullcontext()) @@ -2386,7 +2385,7 @@ def test_pinned_transfer_exits_process_on_failure(monkeypatch): torch.empty(1), ) ] - transfer.status = TransferStatus.RUNNING + transfer.status = "running" logged_messages = [] exit_codes = [] @@ -2424,7 +2423,7 @@ def synchronize(self): torch.tensor([3.0]), ) ] - transfer.status = TransferStatus.IDLE + transfer.status = "idle" transfer._transfer_thread = threading.Thread(target=transfer._run_transfer, daemon=True) receives = [] monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) @@ -2437,9 +2436,9 @@ def synchronize(self): assert not transfer.is_finished() transfer.start() deadline = time.monotonic() + 2 - while transfer.status is TransferStatus.RUNNING and time.monotonic() < deadline: + while transfer.status == "running" and time.monotonic() < deadline: time.sleep(0.001) - assert transfer.status is TransferStatus.SUCCEEDED + assert transfer.status == "succeeded" assert transfer.is_finished() assert torch.equal(transfer.tensor_buffers[0].pinned_row, torch.tensor([3.0])) assert len(receives) == 1 From 59245092a0c826979abbf93d7464ae8dae1fdbfd Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sat, 26 Sep 2026 14:02:12 +0000 Subject: [PATCH 70/72] refactor(eplb): unify async task lifecycle --- docs/CN/source/framework/eplb.md | 4 +- ...t_transfer.py => async_expert_transfer.py} | 96 ++++++++----------- .../eplb/async_placement_plan_task.py | 31 ++++++ .../mode_backend/eplb/async_task.py | 43 +++++++++ .../eplb/async_transfer_planner.py | 62 ++++-------- .../mode_backend/eplb/placement_plan_task.py | 55 ----------- .../mode_backend/eplb/runtime_manager.py | 6 +- unit_tests/common/fused_moe/test_eplb.py | 33 +++---- .../fused_moe/test_eplb_transfer_gpu.py | 7 +- 9 files changed, 157 insertions(+), 180 deletions(-) rename lightllm/server/router/model_infer/mode_backend/eplb/{expert_transfer.py => async_expert_transfer.py} (86%) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/async_task.py delete mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py diff --git a/docs/CN/source/framework/eplb.md b/docs/CN/source/framework/eplb.md index 3161f1e0a1..675e8e0416 100644 --- a/docs/CN/source/framework/eplb.md +++ b/docs/CN/source/framework/eplb.md @@ -93,8 +93,10 @@ prefill_route_counter(最近 24 次 prefill 采样) | `eplb/placement/planner.py` | 布局规划器抽象接口 | | `eplb/placement/factory.py` | 根据 `eplb_plan_mode` 创建具体规划器 | | `eplb/placement/greedy.py` | 默认的贪心布局算法 | +| `eplb/async_task.py` | 统一后台线程任务的启动、完成与异常处理 | +| `eplb/async_placement_plan_task.py` | 在后台根据负载生成目标专家布局 | | `eplb/async_transfer_planner.py` | 在后台生成跨层传输批次 | -| `eplb/expert_transfer.py` | 规划槽位依赖并执行专家权重传输 | +| `eplb/async_expert_transfer.py` | 规划槽位依赖并在后台执行专家权重传输 | | `eplb/runtime_manager.py` | 驱动状态机,协调采集、规划、传输和提交 | ## 4. 初始化布局与权重加载 diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_expert_transfer.py similarity index 86% rename from lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py rename to lightllm/server/router/model_infer/mode_backend/eplb/async_expert_transfer.py index 2943b4fad6..90e9f40acd 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/expert_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_expert_transfer.py @@ -1,7 +1,5 @@ -"""EPLB 专家权重的逐层迁移。""" +"""EPLB 专家权重的异步逐层迁移。""" -import os -import threading import zlib from dataclasses import dataclass from typing import List, Sequence @@ -10,12 +8,10 @@ import torch.distributed as dist from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight -from lightllm.utils.log_utils import init_logger +from .async_task import EPLBAsyncTask from .eplb_utils import NamedTensor, extract_eplb_expert_tensors -logger = init_logger(__name__) - @dataclass(frozen=True) class ExpertTensorBuffer: @@ -43,7 +39,7 @@ class EPLBTransferInfo: dest_local_expert_index: int -class PinnedMemoryEPLBTransfer: +class PinnedMemoryEPLBTransfer(EPLBAsyncTask): """在后台线程中传输一个逻辑专家的全部权重张量。 只有源 rank 和目标 rank 参与 Gloo 点对点通信,具体的数据路径是: @@ -102,61 +98,45 @@ def __init__( ) self._device_to_host_stream: torch.cuda.Stream = torch.cuda.Stream(device=self._device) - self.status = "idle" - self._transfer_thread: threading.Thread = threading.Thread( - target=self._run_transfer, - name=f"eplb-transfer-layer-{transfer_info.layer_index}-expert-{transfer_info.source_logical_expert_id}", - daemon=True, + super().__init__( + thread_name=( + f"eplb-transfer-layer-{transfer_info.layer_index}-expert-{transfer_info.source_logical_expert_id}" + ) ) - def start(self) -> None: - """启动构造函数中 transfer_info 描述的异步传输。""" - assert self.status == "idle", "EPLB transfer has already been started" - self.status = "running" - self._transfer_thread.start() - - def is_finished(self) -> bool: - """返回后台传输是否已经成功完成。""" - return self.status == "succeeded" - - def _run_transfer(self) -> None: + def execute(self) -> None: """把指定专家的全部权重行传输到各 rank 的 pinned memory。""" - try: - transfer_info: EPLBTransferInfo = self.transfer_info + transfer_info: EPLBTransferInfo = self.transfer_info + if self._is_source_rank: + torch.cuda.set_device(self._device) + with torch.cuda.stream(self._device_to_host_stream): + for tensor_buffer in self.tensor_buffers: + tensor_buffer.pinned_row.copy_( + tensor_buffer.live_tensor[transfer_info.source_local_expert_index], + non_blocking=True, + ) + # Gloo 读取 pinned row 前,源 rank 必须等待 GPU -> CPU 拷贝完成。 + self._device_to_host_stream.synchronize() + + if transfer_info.source_rank != transfer_info.dest_rank: if self._is_source_rank: - torch.cuda.set_device(self._device) - with torch.cuda.stream(self._device_to_host_stream): - for tensor_buffer in self.tensor_buffers: - tensor_buffer.pinned_row.copy_( - tensor_buffer.live_tensor[transfer_info.source_local_expert_index], - non_blocking=True, - ) - # Gloo 读取 pinned row 前,源 rank 必须等待 GPU -> CPU 拷贝完成。 - self._device_to_host_stream.synchronize() - - if transfer_info.source_rank != transfer_info.dest_rank: - if self._is_source_rank: - for tensor_buffer in self.tensor_buffers: - message_tag = self._build_p2p_message_tag(tensor_buffer.name) - dist.send( - tensor_buffer.pinned_row, - dst=transfer_info.dest_rank, - group=self._p2p_group, - tag=message_tag, - ) - elif self._is_destination_rank: - for tensor_buffer in self.tensor_buffers: - message_tag = self._build_p2p_message_tag(tensor_buffer.name) - dist.recv( - tensor_buffer.pinned_row, - src=transfer_info.source_rank, - group=self._p2p_group, - tag=message_tag, - ) - self.status = "succeeded" - except BaseException: - logger.exception("EPLB transfer failed") - os._exit(1) + for tensor_buffer in self.tensor_buffers: + message_tag = self._build_p2p_message_tag(tensor_buffer.name) + dist.send( + tensor_buffer.pinned_row, + dst=transfer_info.dest_rank, + group=self._p2p_group, + tag=message_tag, + ) + elif self._is_destination_rank: + for tensor_buffer in self.tensor_buffers: + message_tag = self._build_p2p_message_tag(tensor_buffer.name) + dist.recv( + tensor_buffer.pinned_row, + src=transfer_info.source_rank, + group=self._p2p_group, + tag=message_tag, + ) def _build_p2p_message_tag(self, tensor_name: str) -> int: """为当前专家张量生成 source 和 destination 一致的 Gloo 整数 tag。 diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py new file mode 100644 index 0000000000..440dc6f5c1 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py @@ -0,0 +1,31 @@ +"""EPLB 专家布局的异步规划任务。""" + +from typing import Optional + +import torch + +from .async_task import EPLBAsyncTask +from .placement import EPLBPlanner, ExpertPlacement + + +class EPLBPlanTask(EPLBAsyncTask): + """在后台线程中根据全局专家负载生成新布局。""" + + def __init__( + self, + planner: EPLBPlanner, + global_load: torch.Tensor, + current_placement: ExpertPlacement, + ) -> None: + self.planner = planner + self.global_load = global_load + self.current_placement = current_placement + self.result: Optional[ExpertPlacement] = None + super().__init__(thread_name="eplb-plan") + + def execute(self) -> None: + """根据全局 logical expert 负载生成目标布局。""" + self.result = self.planner.plan( + self.global_load.tolist(), + self.current_placement, + ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_task.py new file mode 100644 index 0000000000..b7c901180f --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_task.py @@ -0,0 +1,43 @@ +"""EPLB 后台线程任务的公共生命周期。""" + +import os +import threading + +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +class EPLBAsyncTask: + """提供单次后台任务统一的启动、完成和失败处理。""" + + def __init__(self, *, thread_name: str) -> None: + self.status = "idle" + self._thread = threading.Thread( + target=self._run, + name=thread_name, + daemon=True, + ) + + def start(self) -> None: + """启动后台任务;同一个任务对象只能启动一次。""" + assert self.status == "idle", f"{type(self).__name__} has already been started" + self.status = "running" + self._thread.start() + + def is_finished(self) -> bool: + """返回后台任务是否已经成功完成。""" + return self.status == "succeeded" + + def _run(self) -> None: + try: + self.execute() + except BaseException: + logger.exception(f"{type(self).__name__} failed") + os._exit(1) + else: + self.status = "succeeded" + + def execute(self) -> None: + """执行子类定义的具体后台任务。""" + raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py index e5dafaa560..5b8c1fad0d 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py @@ -1,18 +1,13 @@ """EPLB 专家传输计划的异步生成器。""" -import os -import threading from typing import List, Optional -from lightllm.utils.log_utils import init_logger - -from .expert_transfer import EPLBTransferInfo, build_transfer_plan +from .async_task import EPLBAsyncTask +from .async_expert_transfer import EPLBTransferInfo, build_transfer_plan from .placement import ExpertPlacement -logger = init_logger(__name__) - -class EPLBTransferPlanner: +class EPLBTransferPlanner(EPLBAsyncTask): """在后台线程中逐层生成并按 layer 顺序拼接专家传输批次。 输入布局在规划期间保持只读。每层独立调用 ``build_transfer_plan``,因此 @@ -31,39 +26,20 @@ def __init__( self.target_placement = target_placement self.num_logical_experts = num_logical_experts self.world_size = world_size - self.status = "idle" self.result: Optional[List[List[EPLBTransferInfo]]] = None - self._thread = threading.Thread( - target=self._run, - name="eplb-transfer-plan", - daemon=True, - ) - - def start(self) -> None: - """启动异步传输规划。""" - assert self.status == "idle", "EPLB transfer planner has already been started" - self.status = "running" - self._thread.start() - - def is_finished(self) -> bool: - """返回全部层的传输批次是否已经生成。""" - return self.status == "succeeded" - - def _run(self) -> None: - try: - transfer_batches: List[List[EPLBTransferInfo]] = [] - layer_placements = zip(self.current_placement, self.target_placement) - for layer_index, (current_layer, target_layer) in enumerate(layer_placements): - layer_transfer_batches = build_transfer_plan( - current_layer, - target_layer, - layer_index, - self.num_logical_experts, - self.world_size, - ) - transfer_batches.extend(layer_transfer_batches) - self.result = transfer_batches - self.status = "succeeded" - except BaseException: - logger.exception("EPLB transfer planning failed") - os._exit(1) + super().__init__(thread_name="eplb-transfer-plan") + + def execute(self) -> None: + """逐层生成传输批次,并按 layer 顺序保存完整结果。""" + transfer_batches: List[List[EPLBTransferInfo]] = [] + layer_placements = zip(self.current_placement, self.target_placement) + for layer_index, (current_layer, target_layer) in enumerate(layer_placements): + layer_transfer_batches = build_transfer_plan( + current_layer, + target_layer, + layer_index, + self.num_logical_experts, + self.world_size, + ) + transfer_batches.extend(layer_transfer_batches) + self.result = transfer_batches diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py deleted file mode 100644 index b7b06bab27..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement_plan_task.py +++ /dev/null @@ -1,55 +0,0 @@ -"""EPLB 专家布局的异步规划任务。""" - -import os -import threading -from typing import Optional - -import torch - -from lightllm.utils.log_utils import init_logger - -from .placement import EPLBPlanner, ExpertPlacement - -logger = init_logger(__name__) - - -class EPLBPlanTask: - """在后台线程中根据全局专家负载生成新布局。""" - - def __init__( - self, - planner: EPLBPlanner, - global_load: torch.Tensor, - current_placement: ExpertPlacement, - ) -> None: - self.planner = planner - self.global_load = global_load - self.current_placement = current_placement - self.status = "idle" - self.result: Optional[ExpertPlacement] = None - self._thread = threading.Thread( - target=self._run, - name="eplb-plan", - daemon=True, - ) - - def start(self) -> None: - """启动异步规划。""" - assert self.status == "idle", "EPLB plan task has already been started" - self.status = "running" - self._thread.start() - - def is_finished(self) -> bool: - """返回规划任务是否已经成功完成。""" - return self.status == "succeeded" - - def _run(self) -> None: - try: - self.result = self.planner.plan( - self.global_load.tolist(), - self.current_placement, - ) - self.status = "succeeded" - except BaseException: - logger.exception("EPLB planning failed") - os._exit(1) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 48d185321d..b2a32fd4f7 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -21,11 +21,12 @@ from lightllm.utils.shm_port_args import get_shm_port_args from . import metrics as eplb_metrics -from .async_transfer_planner import EPLBTransferPlanner -from .expert_transfer import ( +from .async_expert_transfer import ( EPLBTransferInfo, PinnedMemoryEPLBTransfer, ) +from .async_placement_plan_task import EPLBPlanTask +from .async_transfer_planner import EPLBTransferPlanner from .placement import ( EPLBPlanner, ExpertPlacement, @@ -33,7 +34,6 @@ create_eplb_planner, save_placement_config, ) -from .placement_plan_task import EPLBPlanTask logger = init_logger(__name__) EPLB_EXPERT_ALIGNMENT = 128 diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 9c12a37026..b033f2becb 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -19,12 +19,13 @@ from lightllm.server.router.model_infer.mode_backend.eplb import ( runtime_manager as manager_module, ) +from lightllm.server.router.model_infer.mode_backend.eplb import async_task as async_task_module from lightllm.server.router.model_infer.mode_backend.eplb import metrics as eplb_metrics from lightllm.server.router.model_infer.mode_backend.eplb import ( - placement_plan_task as plan_module, + async_placement_plan_task as plan_module, ) from lightllm.server.router.model_infer.mode_backend.eplb import ( - expert_transfer as transfer_module, + async_expert_transfer as transfer_module, ) from lightllm.server.router.model_infer.mode_backend.eplb import ( async_transfer_planner as transfer_planner_module, @@ -44,7 +45,7 @@ fused_moe_weight as fused_weight_module, ) from lightllm.server.router.model_infer.mode_backend.eplb.eplb_utils import extract_eplb_expert_tensors -from lightllm.server.router.model_infer.mode_backend.eplb.expert_transfer import ( +from lightllm.server.router.model_infer.mode_backend.eplb.async_expert_transfer import ( EPLBTransferInfo, ExpertTensorBuffer, PinnedMemoryEPLBTransfer, @@ -845,13 +846,13 @@ def fail(_load, _placement): ) exits = [] logs = [] - monkeypatch.setattr(plan_module.os, "_exit", exits.append) - monkeypatch.setattr(plan_module.logger, "exception", logs.append) + monkeypatch.setattr(async_task_module.os, "_exit", exits.append) + monkeypatch.setattr(async_task_module.logger, "exception", logs.append) task._run() assert exits == [1] - assert logs == ["EPLB planning failed"] + assert logs == ["EPLBPlanTask failed"] def test_transfer_planner_combines_all_layer_batches(monkeypatch): @@ -898,13 +899,13 @@ def fail(*_args): exits = [] logs = [] monkeypatch.setattr(transfer_planner_module, "build_transfer_plan", fail) - monkeypatch.setattr(transfer_planner_module.os, "_exit", exits.append) - monkeypatch.setattr(transfer_planner_module.logger, "exception", logs.append) + monkeypatch.setattr(async_task_module.os, "_exit", exits.append) + monkeypatch.setattr(async_task_module.logger, "exception", logs.append) planner._run() assert exits == [1] - assert logs == ["EPLB transfer planning failed"] + assert logs == ["EPLBTransferPlanner failed"] def test_compute_critical_overhead_ratio_estimates_rank_pressure(): @@ -2327,7 +2328,7 @@ def synchronize(self): lambda tensor, dst, group, tag: sends.append((tensor.clone(), dst, group, tag)), ) - transfer._run_transfer() + transfer._run() assert transfer.status == "succeeded" assert len(sends) == 1 @@ -2364,7 +2365,7 @@ def synchronize(self): monkeypatch.setattr(transfer_module.dist, "send", lambda *_args, **_kwargs: p2p_calls.append("send")) monkeypatch.setattr(transfer_module.dist, "recv", lambda *_args, **_kwargs: p2p_calls.append("recv")) - transfer._run_transfer() + transfer._run() assert transfer.is_finished() assert p2p_calls == [] @@ -2394,12 +2395,12 @@ def fail_recv(*_args, **_kwargs): monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) monkeypatch.setattr(transfer_module.dist, "recv", fail_recv) - monkeypatch.setattr(transfer_module.logger, "exception", logged_messages.append) - monkeypatch.setattr(transfer_module.os, "_exit", exit_codes.append) + monkeypatch.setattr(async_task_module.logger, "exception", logged_messages.append) + monkeypatch.setattr(async_task_module.os, "_exit", exit_codes.append) - transfer._run_transfer() + transfer._run() - assert logged_messages == ["EPLB transfer failed"] + assert logged_messages == ["PinnedMemoryEPLBTransfer failed"] assert exit_codes == [1] assert not transfer.is_finished() @@ -2424,7 +2425,7 @@ def synchronize(self): ) ] transfer.status = "idle" - transfer._transfer_thread = threading.Thread(target=transfer._run_transfer, daemon=True) + transfer._thread = threading.Thread(target=transfer._run, daemon=True) receives = [] monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) monkeypatch.setattr( diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 99c21e721e..36825d7e5e 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -11,10 +11,9 @@ import torch.distributed as dist import torch.multiprocessing as mp -from lightllm.server.router.model_infer.mode_backend.eplb.expert_transfer import ( +from lightllm.server.router.model_infer.mode_backend.eplb.async_expert_transfer import ( EPLBTransferInfo, PinnedMemoryEPLBTransfer, - TransferStatus, build_transfer_plan, ) @@ -82,7 +81,7 @@ def _free_port(): def _wait_for_transfer(transfer, control_group): deadline = time.monotonic() + 30 while time.monotonic() < deadline: - ready_count = torch.tensor([int(transfer.status is TransferStatus.SUCCEEDED)], dtype=torch.int32) + ready_count = torch.tensor([int(transfer.is_finished())], dtype=torch.int32) dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=control_group) if int(ready_count.item()) == 1: return @@ -94,7 +93,7 @@ def _wait_for_all_transfers(transfers, control_group): """等待每个 rank 参与的全部并发传输完成。""" deadline = time.monotonic() + 120 while time.monotonic() < deadline: - local_finished = all(transfer.status is TransferStatus.SUCCEEDED for transfer in transfers) + local_finished = all(transfer.is_finished() for transfer in transfers) globally_finished = torch.tensor([int(local_finished)], dtype=torch.int32) dist.all_reduce(globally_finished, op=dist.ReduceOp.MIN, group=control_group) if int(globally_finished.item()) == 1: From 3fa9b50c065b3c7fb2782be93afaa3cd1e180a0b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 27 Sep 2026 01:30:05 +0000 Subject: [PATCH 71/72] feat(eplb): gather raw load samples asynchronously --- docs/CN/source/framework/eplb.md | 49 ++- .../eplb/async_load_gather_task.py | 43 +++ .../eplb/async_placement_plan_task.py | 6 +- .../model_infer/mode_backend/eplb/metrics.py | 2 +- .../mode_backend/eplb/placement/__init__.py | 3 +- .../mode_backend/eplb/placement/greedy.py | 31 +- .../mode_backend/eplb/placement/planner.py | 13 +- .../mode_backend/eplb/placement/types.py | 2 - .../mode_backend/eplb/runtime_manager.py | 131 +++++-- unit_tests/common/fused_moe/test_eplb.py | 360 +++++++++++++----- 10 files changed, 456 insertions(+), 184 deletions(-) create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb/async_load_gather_task.py diff --git a/docs/CN/source/framework/eplb.md b/docs/CN/source/framework/eplb.md index 675e8e0416..1dfe27d3fc 100644 --- a/docs/CN/source/framework/eplb.md +++ b/docs/CN/source/framework/eplb.md @@ -94,6 +94,7 @@ prefill_route_counter(最近 24 次 prefill 采样) | `eplb/placement/factory.py` | 根据 `eplb_plan_mode` 创建具体规划器 | | `eplb/placement/greedy.py` | 默认的贪心布局算法 | | `eplb/async_task.py` | 统一后台线程任务的启动、完成与异常处理 | +| `eplb/async_load_gather_task.py` | 在独立 Gloo 通信组中后台汇集逐 rank、逐 sample 的原始负载 | | `eplb/async_placement_plan_task.py` | 在后台根据负载生成目标专家布局 | | `eplb/async_transfer_planner.py` | 在后台生成跨层传输批次 | | `eplb/async_expert_transfer.py` | 规划槽位依赖并在后台执行专家权重传输 | @@ -249,7 +250,9 @@ ready 发布使用 `release`、等待方使用 `acquire`,最后完成信号使 ### 5.7 manager 聚合与重置 -manager 在安全推理边界一次性堆叠各层 `[24, E]` 环形缓冲区,在 GPU 上沿 sample 维求和后复制到 CPU,得到 planner 使用的 `[layer, logical_expert]` 负载。如果样本量不足或规划结果未改变布局,不主动清空缓冲区;后续 prefill 会继续写入,并在容量用满后滚动覆盖最旧行。 +manager 在安全推理边界一次性堆叠各层 `[24, E]` 环形缓冲区,并完整复制为 `[layer, sample, logical_expert]` CPU 快照。rank 0 先沿 sample 维聚合本地负载并判断样本量,再向所有 rank 广播是否继续规划;样本不足时直接返回采集状态,不发起大块通信。样本充足时,后台任务使用独立的 Gloo 通信组执行 all-gather,得到 `[rank, layer, sample, logical_expert]`,并把这个四维 Tensor 直接交给 planner。独立通信组使长时间运行的后台 all-gather 不会打乱主线程控制 collective 的调用顺序。 + +如果样本量不足或规划结果未改变布局,不主动清空缓冲区;后续 prefill 会继续写入,并在容量用满后滚动覆盖最旧行。 初始化、成功切换到新布局,以及达到重排次数上限后开始下一轮指标窗口时,manager 会同时清零 `prefill_route_counter` 和 `prefill_route_sample_index`。清零提交到 overlap stream,自然排在此前 forward 之后、后续 forward 之前,不需要额外的全设备同步。 @@ -263,11 +266,16 @@ manager 在安全推理边界一次性堆叠各层 `[24, E]` 环形缓冲区, | v [EVALUATING] - CPU 快照、指标上报、样本量与重排次数检查 + CPU 原始样本快照、指标上报;rank 0 判断样本量并广播结果 + 样本充足且次数允许时启动后台 load all-gather + | + v +[WAIT_LOAD_GATHER_FINISHED] + 等待各 rank 汇集完成,保留 planner 需要的原始四维输入 | v [PLAN_PLACEMENT] - 汇集全局负载,rank 0 启动后台布局规划 + rank 0 启动后台布局规划 | v [WAIT_PLAN_PLACEMENT_FINISHED] @@ -292,7 +300,7 @@ manager 在安全推理边界一次性堆叠各层 `[24, E]` 环形缓冲区, ```text EVALUATING - |-- 平均 token 数不足 ----------> 保留环形窗口,继续滚动采样 + |-- rank 0 平均 token 数不足 ----> 保留环形窗口,继续滚动采样 `-- 达到重排次数上限 ----------> 清空采样,只做周期性指标上报 WAIT_PLAN_PLACEMENT_FINISHED @@ -301,7 +309,7 @@ WAIT_PLAN_PLACEMENT_FINISHED 默认每 20 个采样 step 评估一次,可以通过环境变量 `LIGHTLLM_EPLB_STEP_INTERVAL` 调整。该值必须大于 0。 -只有当整个 world 的平均样本量达到每个“层 × 逻辑专家”128 个 token 时才开始规划。样本不足不会清空环形缓冲区,低流量服务可以跨多个评估周期继续采样;缓冲区写满后只保留最近 24 次 prefill dispatch。 +只有当 rank 0 的平均样本量达到每个“层 × 逻辑专家”128 个 token 时才开始规划。各 rank 的路由分布高度相似,因此 rank 0 足以作为是否值得发起全量通信的低成本判断。样本不足不会清空环形缓冲区,低流量服务可以跨多个评估周期继续采样;缓冲区写满后只保留最近 24 次 prefill dispatch。 ## 7. 专家分布分析 @@ -310,7 +318,7 @@ WAIT_PLAN_PLACEMENT_FINISHED ### 7.1 各 rank 分布相似性(实测) -在 `EPLBManager` 全局汇总负载(`_step_plan_placement` 的 `all_gather` 之后)时, +在 `EPLBManager` 全局汇总负载(后台原始 load `all_gather` 完成之后)时, 把各 rank 的 `[layer][logical_expert]` 负载逐层归一化为概率分布,以 rank0 的分布为基准, 与其余 rank 逐层计算 cosine 相似度。观测窗口为 warmup 阶段一次完整采样 (58 个 MoE 层 × 7 对,共 406 对): @@ -334,21 +342,21 @@ eplb load prob cosine summary mean_by_rank(0..7): 1.0000 0.9773 0.9711 0.9745 0. DP 随机分流下每个 rank 的路由统计都是对全局路由分布的无偏采样, 分布形状(倾斜度、热点名单、长尾形态)在 rank 之间同源。 -### 7.2 设计依据:分布规划只使用 rank0 的数据 +### 7.2 设计选择:异步汇集所有 rank 的原始样本 -上述相似性是后续设计中**分布规划只使用 rank0 负载统计**的合理性依据: +各 rank 的负载分布虽然高度相似,但单 rank 仍带有可观测的采样噪声。当前实现汇集所有 +rank 的 `[layer, sample, logical_expert]` 原始快照,并在通信完成后统一求和: -1. **代表性**:各 rank 分布形状同源且高度一致(cosine 中位 0.97,浅层几乎重合), - rank0 的逐层概率分布与全局聚合分布在形状上等价, - 以 rank0 为样本做倾斜度分档、热点估计、副本预算分配(注水)不会产生系统性偏差。 -2. **成本**:规划无需等待全 rank 负载汇聚即可获得分布形状, - 采集与规划的关键路径缩短到单 rank 统计,状态机的全局同步点相应减少。 -3. **误差边界**:单 rank 估计相对全局的偏差上界为观测到的采样噪声 - (cosine ≥ 0.94),配合负载估计的平滑处理与迁移增益门槛, - 不会因单 rank 采样波动触发错误的副本迁移决策。 +1. **降低采样噪声**:规划器使用整个 world 的累计流量,热点排序和副本预算不依赖某个 + rank 的随机流量分片; +2. **保留分析信息**:通信结果在聚合前保留 rank 和 sample 维,后续可以直接增加跨 rank + 差异或采样稳定性指标,不需要重新设计采集路径; +3. **隔离关键路径**:all-gather 在后台线程和专用 Gloo 通信组中运行,主推理线程只在 + `WAIT_LOAD_GATHER_FINISHED` 状态轮询完成标记,不会被大块负载通信直接阻塞。 -因此布局规划中所有"分布形状"相关的决策(概率分布、倾斜度、热点排序) -均以 rank0 的 logical expert 负载统计为准。 +planner 接口直接接收 `[rank, layer, sample, logical_expert]` CPU Tensor。当前 Greedy +实现进入算法主体前沿 rank 和 sample 维求和为 `[layer, logical_expert]`,再转换成嵌套 +list;因此原始维度在 planner 边界仍然可用,而后续贪心逻辑保持简单的纯 Python 实现。 ## 8. 布局规划 @@ -357,6 +365,7 @@ DP 随机分流下每个 rank 的路由统计都是对全局路由分布的无 所有布局算法实现统一的 `EPLBPlanner.plan(logical_expert_load, current_placement)` 接口,返回: ```text +logical_expert_load: CPU Tensor[rank, layer, sample, logical_expert] [layer][rank][local physical slot] -> logical expert ID ``` @@ -455,14 +464,14 @@ manager 先对每个有效层计算 `max(expert_load) / mean(expert_load)`,过 这三个值都以 `1` 表示完全均衡。例如 P50 为 `1.8`,表示中位层最热 logical expert 的 token 数是该层专家平均值的 1.8 倍。 -当全局样本量达到规划阈值且 planner 产生目标布局后,rank 0 使用 planner 实际消费的同一份 `global_load` 上报: +当样本量达到规划阈值且 planner 产生目标布局后,rank 0 固定选取原始四维快照的第 0 个 sample 行,并仅沿 rank 维汇总为 `[layer, logical_expert]` 负载后上报: ```text lightllm_prefill_ep_compute_critical_overhead_ratio_before_rebalance lightllm_prefill_ep_compute_critical_overhead_ratio_after_rebalance ``` -这两个指标参考 `eplb2` 的关键路径计算开销定义,但不维护独立的 compute counter、后台 monitor 线程和额外通信组。manager 假设同一 logical expert 的流量由 hash 均匀分配给全部 physical 副本,并按 128 token 对每个副本的估算负载向上对齐。`before_rebalance` 使用当前布局,`after_rebalance` 使用 planner 给出的目标布局;二者的输入负载完全相同,可以直接衡量预计的重排收益。如果 planner 判断布局无需改变,两个值应相同。 +这两个指标参考 `eplb2` 的关键路径计算开销定义,但不维护独立的 compute counter、后台 monitor 线程和额外通信组。manager 假设同一 logical expert 的流量由 hash 均匀分配给全部 physical 副本,并按 128 token 对每个副本的估算负载向上对齐。`before_rebalance` 使用当前布局,`after_rebalance` 使用 planner 给出的目标布局;二者的输入负载完全相同,可以直接衡量预计的重排收益。不会沿 sample 维累加,因为不同 sample 行来自不同 prefill 批次,累加后并不对应任何一次真实计算。如果 planner 判断布局无需改变,两个值应相同。 每层先计算最繁忙 rank 相对平均 rank 的额外负载,最后跨层汇总: diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_load_gather_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_load_gather_task.py new file mode 100644 index 0000000000..68e3d57a44 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_load_gather_task.py @@ -0,0 +1,43 @@ +"""EPLB 原始路由负载的异步汇集任务。""" + +from typing import Optional + +import torch +import torch.distributed as dist + +from .async_task import EPLBAsyncTask + + +class EPLBLoadGatherTask(EPLBAsyncTask): + """在独立 Gloo 通信组中汇集每个 rank 的逐样本路由负载。""" + + def __init__( + self, + *, + local_load: torch.Tensor, + load_gather_group: dist.ProcessGroup, + world_size: int, + ) -> None: + assert local_load.device.type == "cpu" + assert local_load.ndim == 3 + assert world_size > 0 + self.local_load = local_load.contiguous() + self.load_gather_group = load_gather_group + self.world_size = world_size + self.result: Optional[torch.Tensor] = None + super().__init__(thread_name="eplb-load-gather") + + def execute(self) -> None: + """生成 ``[rank, layer, sample, logical_expert]`` 的连续结果。""" + gathered_load = torch.empty( + (self.world_size, *self.local_load.shape), + dtype=self.local_load.dtype, + device=self.local_load.device, + ) + load_by_rank = list(gathered_load.unbind(dim=0)) + dist.all_gather( + load_by_rank, + self.local_load, + group=self.load_gather_group, + ) + self.result = gathered_load diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py index 440dc6f5c1..f05ebf6ce0 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py @@ -14,11 +14,11 @@ class EPLBPlanTask(EPLBAsyncTask): def __init__( self, planner: EPLBPlanner, - global_load: torch.Tensor, + logical_expert_load: torch.Tensor, current_placement: ExpertPlacement, ) -> None: self.planner = planner - self.global_load = global_load + self.logical_expert_load = logical_expert_load self.current_placement = current_placement self.result: Optional[ExpertPlacement] = None super().__init__(thread_name="eplb-plan") @@ -26,6 +26,6 @@ def __init__( def execute(self) -> None: """根据全局 logical expert 负载生成目标布局。""" self.result = self.planner.plan( - self.global_load.tolist(), + self.logical_expert_load, self.current_placement, ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py b/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py index 2323c3273c..745cb9263e 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py @@ -120,7 +120,7 @@ def publish_rebalance_compute_metrics( target_placement: ExpertPlacement, expert_alignment: int, ) -> None: - """使用同一份 planner 输入上报重排前后的关键路径开销。""" + """使用同一个 prefill 样本上报重排前后的关键路径开销。""" before_rebalance_ratio = compute_critical_overhead_ratio( logical_expert_load=global_load, placement=current_placement, diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py index 25df90a986..4d4d664fa0 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/__init__.py @@ -1,6 +1,6 @@ """Expert placement construction, routing metadata, and planning APIs.""" -from .types import ExpertPlacement, LayerPlacement, LogicalExpertLoad, LogicalToPhysicalMap +from .types import ExpertPlacement, LayerPlacement, LogicalToPhysicalMap from .planner import EPLBPlanner from .initial import build_initial_local_expert_ids from .routing import build_logical_to_physical_map @@ -13,7 +13,6 @@ "ExpertPlacement", "GreedyEPLBPlanner", "LayerPlacement", - "LogicalExpertLoad", "LogicalToPhysicalMap", "build_initial_local_expert_ids", "build_logical_to_physical_map", diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py index 3a68ba2dc3..91fbd5b404 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py @@ -1,15 +1,18 @@ -"""使用纯 Python 实现贪心 EPLB 专家布局规划。 +"""使用 CPU Tensor 输入和纯 Python 核心逻辑实现贪心 EPLB 布局规划。 -规划器有意使用嵌套 list,而不是 Tensor。Tensor 转换仅发生在 manager 的 -分布式通信和迁移边界;规划模块不依赖 Tensor,更易于阅读、测试和替换算法。 +入口保留 all-gather 产生的 rank、layer、sample 和 logical expert 维度; +Greedy planner 先完成必要的聚合,再转换为嵌套 list。实际贪心分析仍只使用 +Python 数值和容器,更易于阅读、测试和替换算法。 """ import heapq from math import ceil from typing import List +import torch + from .planner import EPLBPlanner -from .types import ExpertPlacement, ExpertReplicaGroup, LayerPlacement, LogicalExpertLoad +from .types import ExpertPlacement, ExpertReplicaGroup, LayerPlacement class GreedyEPLBPlanner(EPLBPlanner): @@ -145,13 +148,21 @@ def __init__( def plan( self, - logical_expert_load: LogicalExpertLoad, + logical_expert_load: torch.Tensor, current_placement: ExpertPlacement, ) -> ExpertPlacement: - """逐层规划专家布局,再组合成完整的多层布局。""" - # 先把外部输入转换成规划器内部统一使用的 Python 数值类型,并一次性 - # 校验所有层的形状和布局约束。后续每层规划之间没有共享的可变状态。 - load = [[float(value) for value in layer] for layer in logical_expert_load] + """聚合全局逐样本负载,逐层规划并组合成完整的多层布局。""" + # logical_expert_load: [rank, layer, sample, logical_expert] CPU Tensor。 + # Greedy 算法只需要整个采样窗口内每层各 logical expert 的累计负载, + # 因此沿 rank 和 sample 维求和为 [layer, logical_expert],再转成 list + # 进入后续纯 Python 分析逻辑。 + assert logical_expert_load.device.type == "cpu" + assert logical_expert_load.ndim == 4 + assert logical_expert_load.shape[0] == self.world_size + load = logical_expert_load.sum(dim=(0, 2)).to(torch.float64).tolist() + + # 一次性校验所有层的形状和布局约束。后续每层规划之间没有共享的 + # 可变状态。 current = [[[int(expert) for expert in rank] for rank in layer] for layer in current_placement] self._validate_inputs(load, current) @@ -342,7 +353,7 @@ def _reuse_current_slots( def _validate_inputs( self, - logical_expert_load: LogicalExpertLoad, + logical_expert_load: List[List[float]], placement: ExpertPlacement, ) -> None: """拒绝会导致逐层规划静默截断的输入。""" diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py index 866fe8f4aa..391bc53bd5 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py @@ -2,7 +2,9 @@ from abc import ABC, abstractmethod -from .types import ExpertPlacement, LogicalExpertLoad +import torch + +from .types import ExpertPlacement class EPLBPlanner(ABC): @@ -11,7 +13,12 @@ class EPLBPlanner(ABC): @abstractmethod def plan( self, - logical_expert_load: LogicalExpertLoad, + logical_expert_load: torch.Tensor, current_placement: ExpertPlacement, ) -> ExpertPlacement: - """返回 ``[layer][rank][local physical expert]`` 专家布局。""" + """根据 CPU 负载生成 ``[layer][rank][local physical expert]`` 布局。 + + ``logical_expert_load`` 的 shape 为 + ``[rank, layer, sample, logical_expert]``。具体 planner 决定如何聚合 + rank 和 sample 维度。 + """ diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/types.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/types.py index b482f4bcdd..f6534ff85f 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/types.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/types.py @@ -3,8 +3,6 @@ from typing import List, Tuple -# [layer][logical expert] -LogicalExpertLoad = List[List[float]] # [rank][local physical expert] -> logical expert LayerPlacement = List[List[int]] # [layer][rank][local physical expert] -> logical expert diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index b2a32fd4f7..26dd4f938a 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -25,6 +25,7 @@ EPLBTransferInfo, PinnedMemoryEPLBTransfer, ) +from .async_load_gather_task import EPLBLoadGatherTask from .async_placement_plan_task import EPLBPlanTask from .async_transfer_planner import EPLBTransferPlanner from .placement import ( @@ -45,6 +46,7 @@ class EPLBManagerState(Enum): COLLECTING = "collecting" EVALUATING = "evaluating" + WAIT_LOAD_GATHER_FINISHED = "wait_load_gather_finished" PLAN_PLACEMENT = "plan_placement" WAIT_PLAN_PLACEMENT_FINISHED = "wait_plan_placement_finished" PLAN_TRANSFER = "plan_transfer" @@ -64,13 +66,19 @@ class EPLBManager: | 评估周期到达 v [EVALUATING] - 将本地环形样本聚合到 CPU 并上报负载指标;汇总各 rank 的 - token 总数,判断样本量和剩余重排次数。 + 将本地环形样本完整复制到 CPU 并上报负载指标;在独立 Gloo + 通信组中启动逐样本负载的后台 all-gather。rank 0 在通信前判断 + 本地样本量;样本不足或达到重排次数上限时不启动通信。 + | + v + [WAIT_LOAD_GATHER_FINISHED] + 等待所有 rank 完成负载汇集,保留 planner 需要的 + rank 和 sample 原始维度。 | | 样本充足且仍允许重排 v [PLAN_PLACEMENT] - 汇集完整的全局专家负载;rank 0 启动后台布局规划任务。 + rank 0 使用完整的全局专家负载启动后台布局规划任务。 | v [WAIT_PLAN_PLACEMENT_FINISHED] @@ -149,7 +157,10 @@ def __init__( self.max_rebalance_count: int = max_rebalance_count self.completed_rebalance_count: int = 0 - # 分布式通信:控制面与权重传输使用独立的通信组。 + # 分布式通信:后台负载汇集、主线程控制面和后台权重传输分别使用 + # 独立的 Gloo 通信组。负载 all-gather 可能跨越多个 manager step, + # 不能与主线程中按 step 排序的控制 collective 共用同一个 group。 + self.load_gather_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") @@ -195,6 +206,10 @@ def step(self) -> None: self._step_evaluating() return + if self.state is EPLBManagerState.WAIT_LOAD_GATHER_FINISHED: + self._step_wait_load_gather_finished() + return + if self.state is EPLBManagerState.PLAN_PLACEMENT: self._step_plan_placement() return @@ -229,21 +244,22 @@ def _step_collecting(self) -> None: self.state = EPLBManagerState.EVALUATING def _step_evaluating(self) -> None: - """发布本地负载指标,并在次数允许时根据全局样本量决定是否规划。""" + """快照本地逐样本负载,并在独立通信组中启动后台汇集。""" counters = [impl.prefill_route_counter for impl in self._eplb_impls] if any(counter.ndim != 2 or counter.shape[1] != self.num_logical_experts for counter in counters): raise RuntimeError("EPLB prefill route counter shape must be [sample_capacity, num_logical_experts]") if len({counter.shape[0] for counter in counters}) != 1: raise RuntimeError("EPLB prefill route counter capacities must match across layers") - # 在 GPU 上沿 sample 维聚合各层的环形样本,再一次性复制到 CPU, - # 得到 planner 使用的 [layer, logical_expert] 负载,避免逐层发起 - # GPU -> CPU 拷贝。 + # 一次性堆叠各层环形样本并复制到 CPU,保留完整的 + # [layer, sample, logical_expert] 维度。后续后台 all-gather 会继续保留 + # rank 和 sample 维;通信完成后将原始四维快照交给 planner。 # 此处先不清零 GPU 样本:如果样本不足或无需迁移,下一周期会继续 # 滚动覆盖最旧行;达到重排上限或成功切换到新布局后才重置窗口。本轮 # 异步规划使用独立的 CPU 快照,不会与后续的 atomic add 竞争。 - local_load = torch.stack(counters).sum(dim=1).detach().cpu() + local_load_samples = torch.stack(counters).detach().cpu() if self.global_rank == 0: + local_load = local_load_samples.sum(dim=1) eplb_metrics.publish_expert_load_metrics( metric_client=self.metric_client, expert_load=local_load, @@ -258,51 +274,81 @@ def _step_evaluating(self) -> None: self._clear_prefill_route_samples() self.state = EPLBManagerState.COLLECTING else: - # 汇集各 rank 的 token 总数,判断当前统计量是否足以进行布局规划。 - token_count_by_rank = [0] * self.world_size - dist.all_gather_object( - token_count_by_rank, - int(local_load.sum().item()), - group=self.control_group, - ) - average_tokens_per_expert = sum(token_count_by_rank) / local_load.numel() - if average_tokens_per_expert < EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT: - if self.global_rank == 0: + # rank 0 的路由分布足以代表全局分布,因此只使用 rank 0 的本地 + # 样本判断统计量是否充足,再广播布尔决策以保持所有 rank 的状态 + # 转移一致。样本不足时不启动大块原始负载 all-gather。 + has_enough_load = None + if self.global_rank == 0: + average_tokens_per_expert = local_load.sum().item() / local_load.numel() + has_enough_load = average_tokens_per_expert >= EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT + if not has_enough_load: logger.info( "eplb continue collecting average_tokens_per_expert=%.2f threshold=%s", average_tokens_per_expert, EPLB_MIN_AVERAGE_TOKENS_PER_EXPERT, ) - self.state = EPLBManagerState.COLLECTING + + values = [has_enough_load] + dist.broadcast_object_list( + values, + src=0, + group=self.control_group, + ) + has_enough_load = values[0] + assert has_enough_load is not None + + if has_enough_load: + self._load_gather_task = EPLBLoadGatherTask( + local_load=local_load_samples, + load_gather_group=self.load_gather_group, + world_size=self.world_size, + ) + self._load_gather_task.start() + self.state = EPLBManagerState.WAIT_LOAD_GATHER_FINISHED else: - self._local_load = local_load - self.state = EPLBManagerState.PLAN_PLACEMENT + self.state = EPLBManagerState.COLLECTING - def _step_plan_placement(self) -> None: - """汇集全局负载,并由 rank 0 启动异步规划。""" - local_load = self._local_load - del self._local_load - - # 一次分配连续的 [rank][layer][logical_expert] 缓冲区,再沿 rank 维 - # 切出 all_gather 所需的输出 tensor。 - gathered_load = torch.empty( - (self.world_size, *local_load.shape), - dtype=local_load.dtype, - device=local_load.device, + def _step_wait_load_gather_finished(self) -> None: + """等待各 rank 的原始负载汇集完成,并生成 planner 的全局负载。""" + local_finished = self._load_gather_task.is_finished() + finished_by_rank = [False] * self.world_size + dist.all_gather_object( + finished_by_rank, + local_finished, + group=self.control_group, ) - load_by_rank = list(gathered_load.unbind(dim=0)) - dist.all_gather(load_by_rank, local_load, group=self.control_group) - global_load = gathered_load.sum(dim=0) + if not all(finished_by_rank): + return + + gathered_load = self._load_gather_task.result + assert gathered_load is not None + assert gathered_load.ndim == 4 + assert gathered_load.shape[0] == self.world_size + assert gathered_load.shape[1] == len(self._eplb_impls) + assert gathered_load.shape[3] == self.num_logical_experts + del self._load_gather_task + + # gathered_load 保留 [rank, layer, sample, logical_expert] 原始结构并 + # 直接交给 planner。重排计算指标应表示一次真实 prefill 的 + # 关键路径开销,而不是多个不同批次累加后的虚拟大批次。 + # 因此固定取环形缓冲区第 0 个 sample 行,只汇总同一次 + # 分布式 prefill 在各 rank 上的分片,得到 [layer, logical_expert]。 + if self.global_rank == 0: + self._planning_load_samples = gathered_load + metric_load_by_rank = gathered_load[:, :, 0, :] + self._planning_global_load = metric_load_by_rank.sum(dim=0) + self.state = EPLBManagerState.PLAN_PLACEMENT + def _step_plan_placement(self) -> None: + """由 rank 0 使用已汇集的全局负载启动异步规划。""" self.state = EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED if self.global_rank == 0: - # 保留 planner 实际消费的全局负载快照。目标布局产生后,使用同一份 - # 输入分别评估 current/target placement,确保 before/after 可直接比较。 - self._planning_global_load = global_load + # planner 消费保留 rank/sample 维的原始快照;指标使用其中 + # 一个 sample 行,确保 before/after 比较的是同一批负载。 self._plan_task = EPLBPlanTask( - self.planner, - global_load, - self.current_placement, + planner=self.planner, + logical_expert_load=self._planning_load_samples, + current_placement=self.current_placement, ) self._plan_task.start() @@ -327,6 +373,7 @@ def _step_wait_plan_placement_finished(self) -> None: target_placement=placement, expert_alignment=EPLB_EXPERT_ALIGNMENT, ) + del self._planning_load_samples del self._planning_global_load del self._plan_task diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index b033f2becb..8c4c73509b 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -20,6 +20,9 @@ runtime_manager as manager_module, ) from lightllm.server.router.model_infer.mode_backend.eplb import async_task as async_task_module +from lightllm.server.router.model_infer.mode_backend.eplb import ( + async_load_gather_task as load_gather_module, +) from lightllm.server.router.model_infer.mode_backend.eplb import metrics as eplb_metrics from lightllm.server.router.model_infer.mode_backend.eplb import ( async_placement_plan_task as plan_module, @@ -108,6 +111,18 @@ def _initial_expert_placement(num_logical_experts, world_size, num_redundant_exp ) +def _planner_load(load_by_layer, world_size): + """构造 [rank, layer, sample, logical expert] CPU planner 输入。""" + aggregated_load = torch.as_tensor(load_by_layer, dtype=torch.float64) + assert aggregated_load.ndim == 2 + load = torch.zeros( + (world_size, aggregated_load.shape[0], 1, aggregated_load.shape[1]), + dtype=torch.float64, + ) + load[0, :, 0] = aggregated_load + return load + + def _rank_to_logic_expert_ids(redundant_placement, num_logical_experts): num_ranks = len(redundant_placement) num_primary_experts_per_rank = num_logical_experts // num_ranks @@ -339,11 +354,11 @@ def test_eplb_planner_builds_legal_concrete_slot_layout(): expert_alignment=1, ) current = _initial_expert_placement(8, 4, 1).unsqueeze(0).tolist() - load = torch.ones((1, 4, 8), dtype=torch.int64) - load[:, :, 0] = 1000 - load[:, :, 4] = 500 + load = torch.ones((4, 1, 1, 8), dtype=torch.int64) + load[:, :, :, 0] = 1000 + load[:, :, :, 4] = 500 - result = planner.plan(load.sum(dim=1).tolist(), current) + result = planner.plan(load, current) placement = result[0] for row in placement: @@ -359,24 +374,45 @@ def test_eplb_planner_returns_deterministic_layout_for_zero_load_experts(): planner = GreedyEPLBPlanner(2, 1) current = [[[0, 1, 3], [2, 3, 1]]] - result = planner.plan([[0, 0, 0, 0]], current) + result = planner.plan(_planner_load([[0, 0, 0, 0]], world_size=2), current) assert result == [[[0, 1, 3], [2, 0, 1]]] +def test_eplb_planner_aggregates_rank_and_sample_dimensions(): + planner = GreedyEPLBPlanner(2, 1) + current = [[[0, 1, 3], [2, 3, 1]]] + raw_load = torch.tensor( + [ + [[[500, 1, 2, 3], [300, 4, 5, 6]]], + [[[100, 7, 8, 9], [200, 10, 11, 12]]], + ], + dtype=torch.int64, + ) + aggregated_load = raw_load.sum(dim=(0, 2)) + equivalent_raw_load = torch.zeros_like(raw_load) + equivalent_raw_load[0, :, 0] = aggregated_load + + result = planner.plan(raw_load, current) + + assert result == planner.plan(equivalent_raw_load, current) + + def test_eplb_planner_plans_each_layer_independently_then_combines_results(): planner = GreedyEPLBPlanner(2, 1) current_layer = [[0, 1, 3], [2, 3, 1]] current = [[row[:] for row in current_layer], [row[:] for row in current_layer]] - result = planner.plan( + load = _planner_load( [ [1000, 1, 1, 1], [0, 0, 0, 0], ], - current, + world_size=2, ) + result = planner.plan(load, current) + assert result == [ [[0, 1, 3], [2, 0, 1]], [[0, 1, 3], [2, 0, 1]], @@ -387,7 +423,7 @@ def test_eplb_planner_iteratively_places_hot_expert_on_idle_rank(): planner = GreedyEPLBPlanner(2, 1) current = [[[0, 1, 3], [2, 3, 1]]] - result = planner.plan([[1000, 1, 1, 1]], current) + result = planner.plan(_planner_load([[1000, 1, 1, 1]], world_size=2), current) assert result == [[[0, 1, 3], [2, 0, 1]]] @@ -396,7 +432,10 @@ def test_eplb_planner_repeatedly_splits_the_hottest_remaining_expert(): planner = GreedyEPLBPlanner(4, 3) current = _initial_expert_placement(8, 4, 3).unsqueeze(0).tolist() - result = planner.plan([[1000, 900, 800, 700, 1, 1, 1, 1]], current) + result = planner.plan( + _planner_load([[1000, 900, 800, 700, 1, 1, 1, 1]], world_size=4), + current, + ) replica_counts = [sum(expert in row for row in result[0]) for expert in range(8)] assert replica_counts == [4, 4, 4, 4, 1, 1, 1, 1] @@ -527,7 +566,10 @@ def test_eplb_planner_keeps_selected_experts_in_their_current_slots(): planner = GreedyEPLBPlanner(4, 2) current = _initial_expert_placement(8, 4, 2).unsqueeze(0).tolist() - result = planner.plan([[50, 98, 54, 6, 34, 66, 63, 52]], current) + result = planner.plan( + _planner_load([[50, 98, 54, 6, 34, 66, 63, 52]], world_size=4), + current, + ) # 只要专家仍分配在同一个 rank,就保留其原物理槽位。 for current_row, target_row in zip(current[0], result[0]): @@ -542,14 +584,15 @@ def test_eplb_planner_fills_every_rank_with_distinct_nonlocal_experts(): 1, ) current = _initial_expert_placement(16, 4, 1).unsqueeze(0).tolist() - load = torch.randint( + load_by_layer_and_rank = torch.randint( 0, 10000, (1, 4, 16), generator=torch.Generator().manual_seed(2), ) + load = load_by_layer_and_rank.permute(1, 0, 2).unsqueeze(dim=2) - result = planner.plan(load.sum(dim=1).tolist(), current) + result = planner.plan(load, current) assert len(result) == len(current) assert all(len(actual) == len(expected) for actual, expected in zip(result[0], current[0])) @@ -561,9 +604,29 @@ def test_eplb_planner_fills_every_rank_with_distinct_nonlocal_experts(): def test_eplb_planner_supports_multiple_redundant_experts_per_rank(): planner = GreedyEPLBPlanner(4, 3) current = _initial_expert_placement(16, 4, 3).unsqueeze(0).tolist() - load = [ - [22613, 26852, 21852, 23480, 13270, 14695, 28735, 22303, 15324, 19604, 21492, 25458, 14120, 12130, 18620, 22888] - ] + load = _planner_load( + [ + [ + 22613, + 26852, + 21852, + 23480, + 13270, + 14695, + 28735, + 22303, + 15324, + 19604, + 21492, + 25458, + 14120, + 12130, + 18620, + 22888, + ] + ], + world_size=4, + ) result = planner.plan(load, current) @@ -781,7 +844,7 @@ def test_transfer_plan_respects_explicit_target_slots(): } -def test_manager_evaluating_copies_route_counters_to_cpu_without_modifying_them(monkeypatch): +def test_manager_evaluating_starts_raw_load_gather_without_modifying_counters(monkeypatch): counters = [ torch.tensor([[10, 11]], dtype=torch.int64), torch.tensor([[40, 41]], dtype=torch.int64), @@ -800,27 +863,81 @@ def test_manager_evaluating_copies_route_counters_to_cpu_without_modifying_them( manager.num_logical_experts = 2 manager.global_rank = 1 manager.world_size = 1 + manager.load_gather_group = object() manager.control_group = object() manager.max_rebalance_count = -1 manager.completed_rebalance_count = 0 - local_token_counts = [] + tasks = [] - def all_gather_object(output, local_token_count, **_kwargs): - local_token_counts.append(local_token_count) - output[:] = [0] + class LoadGatherTask: + def __init__(self, **kwargs): + self.kwargs = kwargs + self.started = False + tasks.append(self) - monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) + def start(self): + self.started = True + + monkeypatch.setattr(manager_module, "EPLBLoadGatherTask", LoadGatherTask) + monkeypatch.setattr( + manager_module.dist, + "broadcast_object_list", + lambda values, **_kwargs: values.__setitem__(0, True), + ) manager._step_evaluating() - assert local_token_counts == [102] + assert manager.state is manager_module.EPLBManagerState.WAIT_LOAD_GATHER_FINISHED + assert len(tasks) == 1 + assert tasks[0].started + assert tasks[0].kwargs["load_gather_group"] is manager.load_gather_group + assert tasks[0].kwargs["world_size"] == manager.world_size + assert torch.equal( + tasks[0].kwargs["local_load"], + torch.tensor([[[10, 11]], [[40, 41]]], dtype=torch.int64), + ) assert torch.equal(counters[0], torch.tensor([[10, 11]], dtype=torch.int64)) assert torch.equal(counters[1], torch.tensor([[40, 41]], dtype=torch.int64)) +def test_load_gather_task_preserves_rank_layer_and_sample_dimensions(monkeypatch): + local_load = torch.tensor( + [ + [[1, 2], [3, 4]], + [[5, 6], [7, 8]], + ], + dtype=torch.int64, + ) + load_gather_group = object() + calls = [] + + def all_gather(output, local, *, group): + calls.append((local, group)) + output[0].copy_(local) + output[1].copy_(local + 100) + + monkeypatch.setattr(load_gather_module.dist, "all_gather", all_gather) + task = load_gather_module.EPLBLoadGatherTask( + local_load=local_load, + load_gather_group=load_gather_group, + world_size=2, + ) + + task._run() + + assert task.is_finished() + assert len(calls) == 1 + assert torch.equal(calls[0][0], local_load) + assert calls[0][1] is load_gather_group + assert task.result is not None + assert task.result.shape == (2, 2, 2, 2) + assert torch.equal(task.result[0], local_load) + assert torch.equal(task.result[1], local_load + 100) + + def test_manager_delegates_distribution_planning_to_planner_class(): current_placement = [[[0, 1]]] - logical_load = torch.tensor([[10, 20]]) + logical_load = torch.tensor([[[[10, 20]]]]) calls = [] planned_placement = [[[0, 1]]] planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or planned_placement)) @@ -831,7 +948,7 @@ def test_manager_delegates_distribution_planning_to_planner_class(): assert task.status == "succeeded" assert task.result == planned_placement assert len(calls) == 1 - assert calls[0][0] == logical_load.tolist() + assert calls[0][0] is logical_load assert calls[0][1] == current_placement @@ -841,7 +958,7 @@ def fail(_load, _placement): task = plan_module.EPLBPlanTask( SimpleNamespace(plan=fail), - torch.tensor([[10, 20]]), + torch.tensor([[[[10, 20]]]]), [[[1]]], ) exits = [] @@ -1074,51 +1191,63 @@ def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): assert not hasattr(impl, "recording") -def test_manager_evaluation_gathers_token_counts_from_all_ranks(monkeypatch): +def test_manager_wait_load_gather_aggregates_rank_and_sample_dimensions(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "fuse_moe_impl": _test_moe_impl( - eplb=True, - prefill_route_counter=torch.zeros((24, 4), dtype=torch.int64), - num_logical_experts=4, - world_size=1, - ) - }, - )() - ] - manager._eplb_impls = [manager.weights[0].fuse_moe_impl] - manager.global_rank = 2 + manager._eplb_impls = [object()] + manager.global_rank = 0 manager.world_size = 4 - manager.step_interval = 20 manager.num_logical_experts = 4 - manager.num_redundant_experts_per_rank = 1 - manager.current_placement = _initial_expert_placement(4, 4, 1).unsqueeze(0).tolist() manager.control_group = object() - manager.max_rebalance_count = -1 - manager.completed_rebalance_count = 0 - local = torch.full((24, 4), 0, dtype=torch.int64) - local[0].fill_(100) - manager._eplb_impls[0].prefill_route_counter = local - manager.state = manager_module.EPLBManagerState.EVALUATING + manager.state = manager_module.EPLBManagerState.WAIT_LOAD_GATHER_FINISHED + gathered_load = torch.zeros((4, 1, 24, 4), dtype=torch.int64) + gathered_load[0, 0, 0].fill_(100) + gathered_load[1, 0, 1].fill_(100) + gathered_load[2, 0, 0].fill_(50) + manager._load_gather_task = SimpleNamespace( + is_finished=lambda: True, + result=gathered_load, + ) seen = {} - def all_gather_object(output, local_token_count, **kwargs): - seen["local_token_count"] = local_token_count + def all_gather_object(output, local_finished, **kwargs): + seen["local_finished"] = local_finished seen["group"] = kwargs["group"] - # Simulate one other rank contributing the same logical-expert load. - output[:] = [local_token_count, local_token_count, 0, 0] + output[:] = [True] * manager.world_size monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) - manager._step_evaluating() + manager._step_wait_load_gather_finished() assert seen["group"] is manager.control_group - assert seen["local_token_count"] == 400 - assert manager.state is manager_module.EPLBManagerState.COLLECTING + assert seen["local_finished"] is True + assert manager.state is manager_module.EPLBManagerState.PLAN_PLACEMENT + assert manager._planning_load_samples is gathered_load + # metrics 只汇总各 rank 的 sample 0;rank 1 在 sample 1 中的负载 + # 不应累加进来。 + assert torch.equal(manager._planning_global_load, torch.full((1, 4), 150, dtype=torch.int64)) + assert not hasattr(manager, "_load_gather_task") + + +def test_manager_wait_load_gather_does_not_advance_until_every_rank_finishes(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.state = manager_module.EPLBManagerState.WAIT_LOAD_GATHER_FINISHED + manager.world_size = 2 + manager.control_group = object() + load_gather_task = SimpleNamespace( + is_finished=lambda: True, + result=torch.ones((2, 1, 1, 2), dtype=torch.int64), + ) + manager._load_gather_task = load_gather_task + monkeypatch.setattr( + manager_module.dist, + "all_gather_object", + lambda output, _local_finished, **_kwargs: output.__setitem__(slice(None), [True, False]), + ) + + manager._step_wait_load_gather_finished() + + assert manager.state is manager_module.EPLBManagerState.WAIT_LOAD_GATHER_FINISHED + assert manager._load_gather_task is load_gather_task def test_decode_dispatch_uses_physical_ids_and_total_expert_count(monkeypatch): @@ -1869,6 +1998,7 @@ def test_wait_plan_finish_publishes_before_and_after_metrics(monkeypatch): manager.control_group = object() manager.current_placement = [[[0, 1], [2, 3]]] target_placement = [[[0, 2], [1, 3]]] + manager._planning_load_samples = torch.tensor([[[[4, 3, 2, 1]]]], dtype=torch.int64) manager._planning_global_load = torch.tensor([[4, 3, 2, 1]], dtype=torch.int64) manager._plan_task = SimpleNamespace( is_finished=lambda: True, @@ -1890,6 +2020,7 @@ def test_wait_plan_finish_publishes_before_and_after_metrics(monkeypatch): assert len(published) == 1 assert torch.equal(published[0][0], torch.tensor([[4, 3, 2, 1]], dtype=torch.int64)) assert published[0][1] is target_placement + assert not hasattr(manager, "_planning_load_samples") assert not hasattr(manager, "_planning_global_load") assert not hasattr(manager, "_plan_task") @@ -1977,7 +2108,7 @@ def test_manager_step_advances_inflight_transfer(): def test_manager_evaluates_only_after_entering_evaluating_state(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) prefill_route_counter = torch.tensor([[1, 2]], dtype=torch.int64) - local_token_counts = [] + load_gather_tasks = [] manager.state = manager_module.EPLBManagerState.COLLECTING manager.global_rank = 1 manager.steps = 0 @@ -1989,12 +2120,16 @@ def test_manager_evaluates_only_after_entering_evaluating_state(monkeypatch): manager.control_group = object() manager.max_rebalance_count = -1 manager.completed_rebalance_count = 0 - - def all_gather_object(output, local_token_count, **_kwargs): - local_token_counts.append(local_token_count) - output[:] = [0] - - monkeypatch.setattr(manager_module.dist, "all_gather_object", all_gather_object) + monkeypatch.setattr( + manager_module, + "EPLBLoadGatherTask", + lambda **_kwargs: load_gather_tasks.append(True), + ) + monkeypatch.setattr( + manager_module.dist, + "broadcast_object_list", + lambda values, **_kwargs: values.__setitem__(0, False), + ) manager.step() manager.step() @@ -2002,13 +2137,13 @@ def all_gather_object(output, local_token_count, **_kwargs): manager.step() assert manager.state is manager_module.EPLBManagerState.EVALUATING - assert local_token_counts == [] + assert load_gather_tasks == [] assert manager.next_evaluation_step == 6 manager.step() assert manager.state is manager_module.EPLBManagerState.COLLECTING - assert local_token_counts == [3] + assert load_gather_tasks == [] assert manager.next_evaluation_step == 6 assert torch.equal(prefill_route_counter, torch.tensor([[1, 2]], dtype=torch.int64)) @@ -2122,26 +2257,36 @@ def gather_finished(output, local_finished, **_kwargs): def test_manager_evaluation_with_insufficient_tokens_returns_to_collecting(monkeypatch): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.EVALUATING - manager._eplb_impls = [SimpleNamespace(prefill_route_counter=torch.full((1, 4), 255, dtype=torch.int64))] + manager._eplb_impls = [SimpleNamespace(prefill_route_counter=torch.full((1, 4), 100, dtype=torch.int64))] manager.num_logical_experts = 4 manager.steps = 11 manager.step_interval = 20 manager.next_evaluation_step = 31 - manager.global_rank = 1 + manager.global_rank = 0 manager.world_size = 1 manager.control_group = object() + manager.load_gather_group = object() manager.max_rebalance_count = -1 manager.completed_rebalance_count = 0 + manager.metric_client = object() + broadcast_decisions = [] + monkeypatch.setattr(eplb_metrics, "publish_expert_load_metrics", lambda **_kwargs: None) + monkeypatch.setattr( + manager_module, + "EPLBLoadGatherTask", + lambda **_kwargs: pytest.fail("insufficient rank-0 load must skip all-gather"), + ) monkeypatch.setattr( manager_module.dist, - "all_gather_object", - lambda output, local_token_count, **_kwargs: output.__setitem__(slice(None), [local_token_count]), + "broadcast_object_list", + lambda values, **kwargs: broadcast_decisions.append((values[0], kwargs["src"], kwargs["group"])), ) manager.step() assert manager.state is manager_module.EPLBManagerState.COLLECTING assert manager.next_evaluation_step == 31 + assert broadcast_decisions == [(False, 0, manager.control_group)] def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): @@ -2155,23 +2300,34 @@ def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): manager._eplb_impls = [SimpleNamespace(prefill_route_counter=local_load)] manager.world_size = 1 manager.control_group = object() + manager.load_gather_group = object() manager.max_rebalance_count = -1 manager.completed_rebalance_count = 0 + + class LoadGatherTask: + def __init__(self, **kwargs): + self.local_load = kwargs["local_load"] + self.result = self.local_load.unsqueeze(0) + self.started = False + + def start(self): + self.started = True + + def is_finished(self): + return True + + monkeypatch.setattr(manager_module, "EPLBLoadGatherTask", LoadGatherTask) + monkeypatch.setattr(manager_module.dist, "broadcast_object_list", lambda _values, **_kwargs: None) monkeypatch.setattr( manager_module.dist, "all_gather_object", - lambda output, local_token_count, **_kwargs: output.__setitem__(slice(None), [local_token_count]), - ) - monkeypatch.setattr( - manager_module.dist, - "all_gather", - lambda output, local, **_kwargs: output[0].copy_(local), + lambda output, local_finished, **_kwargs: output.__setitem__(slice(None), [local_finished]), ) class PlanTask: - def __init__(self, planner, global_load, current_placement): + def __init__(self, planner, logical_expert_load, current_placement): self.planner = planner - self.global_load = global_load + self.logical_expert_load = logical_expert_load self.current_placement = current_placement self.started = False plan_tasks.append(self) @@ -2191,8 +2347,9 @@ def start(self): manager.step() - assert manager.state is manager_module.EPLBManagerState.PLAN_PLACEMENT - assert torch.equal(manager._local_load, local_load) + assert manager.state is manager_module.EPLBManagerState.WAIT_LOAD_GATHER_FINISHED + assert manager._load_gather_task.started + assert torch.equal(manager._load_gather_task.local_load, local_load.unsqueeze(0)) assert len(published_loads) == 1 assert torch.equal(published_loads[0], local_load) assert plan_tasks == [] @@ -2200,14 +2357,23 @@ def start(self): manager.step() + assert manager.state is manager_module.EPLBManagerState.PLAN_PLACEMENT + planning_load_samples = manager._planning_load_samples + assert torch.equal(planning_load_samples, local_load.unsqueeze(0).unsqueeze(0)) + assert torch.equal(manager._planning_global_load, local_load) + assert not hasattr(manager, "_load_gather_task") + assert plan_tasks == [] + + manager.step() + assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED - assert torch.equal(plan_tasks[0].global_load, local_load) + assert plan_tasks[0].logical_expert_load is planning_load_samples + assert torch.equal(plan_tasks[0].logical_expert_load, local_load.unsqueeze(0).unsqueeze(0)) assert plan_tasks[0].planner is manager.planner assert plan_tasks[0].current_placement is manager.current_placement assert plan_tasks[0].started assert manager._plan_task is plan_tasks[0] assert torch.equal(manager._planning_global_load, local_load) - assert not hasattr(manager, "_local_load") assert len(published_loads) == 1 @@ -2230,9 +2396,9 @@ def test_manager_keeps_reporting_after_reaching_rebalance_limit(monkeypatch): ) manager._clear_prefill_route_samples = lambda: cleared_counters.append(True) monkeypatch.setattr( - manager_module.dist, - "all_gather_object", - lambda *_args, **_kwargs: pytest.fail("rebalancing must stop after reaching the limit"), + manager_module, + "EPLBLoadGatherTask", + lambda **_kwargs: pytest.fail("load gathering must stop after reaching the limit"), ) manager.step() @@ -2243,25 +2409,17 @@ def test_manager_keeps_reporting_after_reaching_rebalance_limit(monkeypatch): assert cleared_counters == [True] -def test_nonzero_rank_waits_without_starting_planner(monkeypatch): +def test_nonzero_rank_enters_plan_wait_without_starting_planner(): manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) manager.state = manager_module.EPLBManagerState.PLAN_PLACEMENT manager.global_rank = 1 - manager.world_size = 2 - manager.control_group = object() - manager._local_load = torch.tensor([[1, 2]], dtype=torch.int64) - - def all_gather(output, local, **_kwargs): - output[0].copy_(local) - output[1].copy_(local) - - monkeypatch.setattr(manager_module.dist, "all_gather", all_gather) manager.step() assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED assert not hasattr(manager, "_plan_task") - assert not hasattr(manager, "_local_load") + assert not hasattr(manager, "_planning_load_samples") + assert not hasattr(manager, "_planning_global_load") def test_manager_planning_without_changes_returns_to_collecting(monkeypatch): @@ -2508,7 +2666,7 @@ def test_manager_initializes_without_transfer_task(monkeypatch): ), }, )() - groups = [object(), object()] + groups = [object(), object(), object()] new_group_calls = [] monkeypatch.setattr(manager_module, "is_sm100_gpu", lambda: False) monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) @@ -2554,9 +2712,9 @@ def all_gather_object(output, local_expert_ids_by_layer, group): assert not hasattr(manager, "_plan_task") assert not hasattr(manager, "pending_transfer_batches") assert manager.state is manager_module.EPLBManagerState.COLLECTING - assert (manager.control_group, manager.transfer_group) == tuple(groups) - assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 2 - assert all_gather_calls == [([[0, 1, 2, 3]], groups[0])] + assert (manager.load_gather_group, manager.control_group, manager.transfer_group) == tuple(groups) + assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 + assert all_gather_calls == [([[0, 1, 2, 3]], groups[1])] assert manager.current_placement == [[[0, 1, 2, 3], [2, 3, 0, 1]]] assert manager.metric_client is metric_client assert metric_client_ports == [1234] From cb02e6dda32ec1e9bb78f6e9070a84a5d2f81ea6 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 27 Sep 2026 03:16:08 +0000 Subject: [PATCH 72/72] refactor(eplb): clarify rebalance state lifecycle --- docs/CN/source/framework/eplb.md | 4 +- .../eplb/async_expert_transfer.py | 1 + .../eplb/async_load_gather_task.py | 21 +- .../eplb/async_placement_plan_task.py | 9 +- .../eplb/async_transfer_planner.py | 11 +- .../model_infer/mode_backend/eplb/metrics.py | 6 +- .../mode_backend/eplb/placement/greedy.py | 22 +- .../mode_backend/eplb/placement/planner.py | 5 +- .../mode_backend/eplb/runtime_manager.py | 206 ++++++++++-------- unit_tests/common/fused_moe/test_eplb.py | 192 +++++++++------- .../fused_moe/test_eplb_transfer_gpu.py | 41 +++- 11 files changed, 306 insertions(+), 212 deletions(-) diff --git a/docs/CN/source/framework/eplb.md b/docs/CN/source/framework/eplb.md index 1dfe27d3fc..23fea0c522 100644 --- a/docs/CN/source/framework/eplb.md +++ b/docs/CN/source/framework/eplb.md @@ -362,10 +362,10 @@ list;因此原始维度在 planner 边界仍然可用,而后续贪心逻辑 ### 8.1 规划器接口与选择 -所有布局算法实现统一的 `EPLBPlanner.plan(logical_expert_load, current_placement)` 接口,返回: +所有布局算法实现统一的 `EPLBPlanner.plan(logical_expert_load_samples, current_placement)` 接口,返回: ```text -logical_expert_load: CPU Tensor[rank, layer, sample, logical_expert] +logical_expert_load_samples: CPU Tensor[rank, layer, sample, logical_expert] [layer][rank][local physical slot] -> logical expert ID ``` diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_expert_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_expert_transfer.py index 90e9f40acd..fe1c79ec4f 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/async_expert_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_expert_transfer.py @@ -60,6 +60,7 @@ class PinnedMemoryEPLBTransfer(EPLBAsyncTask): def __init__( self, + *, weights: Sequence[FusedMoeWeight], transfer_group: dist.ProcessGroup, current_global_rank: int, diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_load_gather_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_load_gather_task.py index 68e3d57a44..48e18efea3 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/async_load_gather_task.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_load_gather_task.py @@ -14,30 +14,29 @@ class EPLBLoadGatherTask(EPLBAsyncTask): def __init__( self, *, - local_load: torch.Tensor, + local_load_samples: torch.Tensor, load_gather_group: dist.ProcessGroup, - world_size: int, ) -> None: - assert local_load.device.type == "cpu" - assert local_load.ndim == 3 - assert world_size > 0 - self.local_load = local_load.contiguous() + assert local_load_samples.device.type == "cpu" + assert local_load_samples.ndim == 3 + self.local_load_samples = local_load_samples.contiguous() self.load_gather_group = load_gather_group - self.world_size = world_size + self.world_size = dist.get_world_size(group=load_gather_group) + assert self.world_size > 0 self.result: Optional[torch.Tensor] = None super().__init__(thread_name="eplb-load-gather") def execute(self) -> None: """生成 ``[rank, layer, sample, logical_expert]`` 的连续结果。""" gathered_load = torch.empty( - (self.world_size, *self.local_load.shape), - dtype=self.local_load.dtype, - device=self.local_load.device, + (self.world_size, *self.local_load_samples.shape), + dtype=self.local_load_samples.dtype, + device=self.local_load_samples.device, ) load_by_rank = list(gathered_load.unbind(dim=0)) dist.all_gather( load_by_rank, - self.local_load, + self.local_load_samples, group=self.load_gather_group, ) self.result = gathered_load diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py index f05ebf6ce0..8d380fd63b 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_placement_plan_task.py @@ -13,12 +13,13 @@ class EPLBPlanTask(EPLBAsyncTask): def __init__( self, + *, planner: EPLBPlanner, - logical_expert_load: torch.Tensor, + logical_expert_load_samples: torch.Tensor, current_placement: ExpertPlacement, ) -> None: self.planner = planner - self.logical_expert_load = logical_expert_load + self.logical_expert_load_samples = logical_expert_load_samples self.current_placement = current_placement self.result: Optional[ExpertPlacement] = None super().__init__(thread_name="eplb-plan") @@ -26,6 +27,6 @@ def __init__( def execute(self) -> None: """根据全局 logical expert 负载生成目标布局。""" self.result = self.planner.plan( - self.logical_expert_load, - self.current_placement, + logical_expert_load_samples=self.logical_expert_load_samples, + current_placement=self.current_placement, ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py index 5b8c1fad0d..c8646dbbe0 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/async_transfer_planner.py @@ -17,6 +17,7 @@ class EPLBTransferPlanner(EPLBAsyncTask): def __init__( self, + *, current_placement: ExpertPlacement, target_placement: ExpertPlacement, num_logical_experts: int, @@ -35,11 +36,11 @@ def execute(self) -> None: layer_placements = zip(self.current_placement, self.target_placement) for layer_index, (current_layer, target_layer) in enumerate(layer_placements): layer_transfer_batches = build_transfer_plan( - current_layer, - target_layer, - layer_index, - self.num_logical_experts, - self.world_size, + current_placement=current_layer, + target_placement=target_layer, + layer_index=layer_index, + num_logical_experts=self.num_logical_experts, + world_size=self.world_size, ) transfer_batches.extend(layer_transfer_batches) self.result = transfer_batches diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py b/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py index 745cb9263e..c94ff2cabb 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/metrics.py @@ -115,19 +115,19 @@ def compute_critical_overhead_ratio( def publish_rebalance_compute_metrics( *, metric_client: MetricClient, - global_load: torch.Tensor, + sample_load: torch.Tensor, current_placement: ExpertPlacement, target_placement: ExpertPlacement, expert_alignment: int, ) -> None: """使用同一个 prefill 样本上报重排前后的关键路径开销。""" before_rebalance_ratio = compute_critical_overhead_ratio( - logical_expert_load=global_load, + logical_expert_load=sample_load, placement=current_placement, expert_alignment=expert_alignment, ) after_rebalance_ratio = compute_critical_overhead_ratio( - logical_expert_load=global_load, + logical_expert_load=sample_load, placement=target_placement, expert_alignment=expert_alignment, ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py index 91fbd5b404..5796c5b8ea 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/greedy.py @@ -148,27 +148,31 @@ def __init__( def plan( self, - logical_expert_load: torch.Tensor, + *, + logical_expert_load_samples: torch.Tensor, current_placement: ExpertPlacement, ) -> ExpertPlacement: """聚合全局逐样本负载,逐层规划并组合成完整的多层布局。""" - # logical_expert_load: [rank, layer, sample, logical_expert] CPU Tensor。 + # logical_expert_load_samples: [rank, layer, sample, logical_expert] + # CPU Tensor。 # Greedy 算法只需要整个采样窗口内每层各 logical expert 的累计负载, # 因此沿 rank 和 sample 维求和为 [layer, logical_expert],再转成 list # 进入后续纯 Python 分析逻辑。 - assert logical_expert_load.device.type == "cpu" - assert logical_expert_load.ndim == 4 - assert logical_expert_load.shape[0] == self.world_size - load = logical_expert_load.sum(dim=(0, 2)).to(torch.float64).tolist() + assert logical_expert_load_samples.device.type == "cpu" + assert logical_expert_load_samples.ndim == 4 + assert logical_expert_load_samples.shape[0] == self.world_size + aggregated_load = logical_expert_load_samples.sum(dim=(0, 2)).to(torch.float64).tolist() # 一次性校验所有层的形状和布局约束。后续每层规划之间没有共享的 # 可变状态。 - current = [[[int(expert) for expert in rank] for rank in layer] for layer in current_placement] - self._validate_inputs(load, current) + self._validate_inputs(aggregated_load, current_placement) # 每层只依赖自己的逻辑专家负载和当前布局。先完成单层规划,再将结果 # 按原 layer 顺序组合,避免多层候选和负载数据交叉索引。 - return [self._plan_layer(layer_load, current_layer) for layer_load, current_layer in zip(load, current)] + return [ + self._plan_layer(layer_load, current_layer) + for layer_load, current_layer in zip(aggregated_load, current_placement) + ] def _plan_layer( self, diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py b/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py index 391bc53bd5..2b9e66da38 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/placement/planner.py @@ -13,12 +13,13 @@ class EPLBPlanner(ABC): @abstractmethod def plan( self, - logical_expert_load: torch.Tensor, + *, + logical_expert_load_samples: torch.Tensor, current_placement: ExpertPlacement, ) -> ExpertPlacement: """根据 CPU 负载生成 ``[layer][rank][local physical expert]`` 布局。 - ``logical_expert_load`` 的 shape 为 + ``logical_expert_load_samples`` 的 shape 为 ``[rank, layer, sample, logical_expert]``。具体 planner 决定如何聚合 rank 和 sample 维度。 """ diff --git a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py index 26dd4f938a..6ff08ebb0f 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb/runtime_manager.py @@ -157,6 +157,19 @@ def __init__( self.max_rebalance_count: int = max_rebalance_count self.completed_rebalance_count: int = 0 + # 一次重排周期中的短期状态。统一初始化为 None,避免各状态 + # 通过 hasattr()/del 隐式定义 EPLBManager 的属性结构。 + self._load_gather_task: Optional[EPLBLoadGatherTask] = None + self._pending_plan_load_samples: Optional[torch.Tensor] = None + self._metric_sample_load: Optional[torch.Tensor] = None + self._plan_task: Optional[EPLBPlanTask] = None + self.target_placement: Optional[ExpertPlacement] = None + self._transfer_planner: Optional[EPLBTransferPlanner] = None + self.pending_transfer_batches: Optional[List[List[EPLBTransferInfo]]] = None + self.rebalance_started_at: Optional[float] = None + self.active_transfer_batch: Optional[List[EPLBTransferInfo]] = None + self.active_transfers: Optional[List[PinnedMemoryEPLBTransfer]] = None + # 分布式通信:后台负载汇集、主线程控制面和后台权重传输分别使用 # 独立的 Gloo 通信组。负载 all-gather 可能跨越多个 manager step, # 不能与主线程中按 step 排序的控制 collective 共用同一个 group。 @@ -299,9 +312,8 @@ def _step_evaluating(self) -> None: if has_enough_load: self._load_gather_task = EPLBLoadGatherTask( - local_load=local_load_samples, + local_load_samples=local_load_samples, load_gather_group=self.load_gather_group, - world_size=self.world_size, ) self._load_gather_task.start() self.state = EPLBManagerState.WAIT_LOAD_GATHER_FINISHED @@ -309,55 +321,57 @@ def _step_evaluating(self) -> None: self.state = EPLBManagerState.COLLECTING def _step_wait_load_gather_finished(self) -> None: - """等待各 rank 的原始负载汇集完成,并生成 planner 的全局负载。""" - local_finished = self._load_gather_task.is_finished() - finished_by_rank = [False] * self.world_size - dist.all_gather_object( - finished_by_rank, - local_finished, - group=self.control_group, - ) - if not all(finished_by_rank): + """等待原始负载汇集完成,并准备 planner 和 metrics 输入。""" + assert self._load_gather_task is not None + if not self._all_ranks_finished(self._load_gather_task.is_finished()): return - gathered_load = self._load_gather_task.result - assert gathered_load is not None - assert gathered_load.ndim == 4 - assert gathered_load.shape[0] == self.world_size - assert gathered_load.shape[1] == len(self._eplb_impls) - assert gathered_load.shape[3] == self.num_logical_experts - del self._load_gather_task + gathered_load_samples = self._load_gather_task.result + assert gathered_load_samples is not None + assert gathered_load_samples.ndim == 4 + assert gathered_load_samples.shape[0] == self.world_size + assert gathered_load_samples.shape[1] == len(self._eplb_impls) + assert gathered_load_samples.shape[3] == self.num_logical_experts + self._load_gather_task = None - # gathered_load 保留 [rank, layer, sample, logical_expert] 原始结构并 + # gathered_load_samples 保留 [rank, layer, sample, logical_expert] + # 原始结构并 # 直接交给 planner。重排计算指标应表示一次真实 prefill 的 # 关键路径开销,而不是多个不同批次累加后的虚拟大批次。 - # 因此固定取环形缓冲区第 0 个 sample 行,只汇总同一次 - # 分布式 prefill 在各 rank 上的分片,得到 [layer, logical_expert]。 + # EP rank 以相同顺序执行 prefill,每次 dispatch 都将同一 sample + # index 推进一次;因此各 rank 的第 0 行属于同一采样位置。 + # metrics 固定取该行,只汇总其在各 rank 上的分片,得到 + # [layer, logical_expert]。 if self.global_rank == 0: - self._planning_load_samples = gathered_load - metric_load_by_rank = gathered_load[:, :, 0, :] - self._planning_global_load = metric_load_by_rank.sum(dim=0) + self._pending_plan_load_samples = gathered_load_samples + metric_sample_load_by_rank = gathered_load_samples[:, :, 0, :] + self._metric_sample_load = metric_sample_load_by_rank.sum(dim=0) self.state = EPLBManagerState.PLAN_PLACEMENT def _step_plan_placement(self) -> None: """由 rank 0 使用已汇集的全局负载启动异步规划。""" self.state = EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED if self.global_rank == 0: + assert self._pending_plan_load_samples is not None # planner 消费保留 rank/sample 维的原始快照;指标使用其中 # 一个 sample 行,确保 before/after 比较的是同一批负载。 self._plan_task = EPLBPlanTask( planner=self.planner, - logical_expert_load=self._planning_load_samples, + logical_expert_load_samples=self._pending_plan_load_samples, current_placement=self.current_placement, ) self._plan_task.start() + # 任务对象已持有 Tensor,manager 不再保留重复引用。 + self._pending_plan_load_samples = None def _step_wait_plan_placement_finished(self) -> None: """等待 rank 0 完成规划并广播目标专家排布。""" placement: Optional[ExpertPlacement] = None - if self.global_rank == 0 and self._plan_task.is_finished(): - placement = self._plan_task.result - assert placement is not None + if self.global_rank == 0: + assert self._plan_task is not None + if self._plan_task.is_finished(): + placement = self._plan_task.result + assert placement is not None values = [placement] dist.broadcast_object_list(values, src=0, group=self.control_group) @@ -366,16 +380,16 @@ def _step_wait_plan_placement_finished(self) -> None: return if self.global_rank == 0: + assert self._metric_sample_load is not None eplb_metrics.publish_rebalance_compute_metrics( metric_client=self.metric_client, - global_load=self._planning_global_load, + sample_load=self._metric_sample_load, current_placement=self.current_placement, target_placement=placement, expert_alignment=EPLB_EXPERT_ALIGNMENT, ) - del self._planning_load_samples - del self._planning_global_load - del self._plan_task + self._metric_sample_load = None + self._plan_task = None if placement == self.current_placement: if self.global_rank == 0: @@ -393,65 +407,67 @@ def _step_plan_transfer(self) -> None: # 所有 rank 使用相同的 current/target placement 独立生成确定性的传输 # 批次,避免广播体积较大的任务列表;耗时的逐层依赖分析放到后台线程, # 当前推理线程从下一次安全边界开始只需轮询完成状态。 + assert self.target_placement is not None self._transfer_planner = EPLBTransferPlanner( - self.current_placement, - self.target_placement, - self.num_logical_experts, - self.world_size, + current_placement=self.current_placement, + target_placement=self.target_placement, + num_logical_experts=self.num_logical_experts, + world_size=self.world_size, ) self._transfer_planner.start() self.state = EPLBManagerState.WAIT_PLAN_TRANSFER_FINISHED def _step_wait_plan_transfer_finished(self) -> None: """等待所有 rank 异步生成相同的传输批次,再统一进入传输状态。""" - local_finished = self._transfer_planner.is_finished() - finished_by_rank = [False] * self.world_size - dist.all_gather_object(finished_by_rank, local_finished, group=self.control_group) + assert self._transfer_planner is not None # 即使本 rank 已经完成,也必须等待其他 rank 的镜像任务列表就绪;否则 # 提前进入 TRANSFERRING 的 rank 可能发起尚无对端参与的点对点传输。 - if not all(finished_by_rank): + if not self._all_ranks_finished(self._transfer_planner.is_finished()): return - else: - pending_transfer_batches = self._transfer_planner.result - assert pending_transfer_batches is not None - self.pending_transfer_batches = pending_transfer_batches - del self._transfer_planner - if not self.pending_transfer_batches: - raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") - - self.state = EPLBManagerState.TRANSFERRING - if self.global_rank == 0: - changed_layer_count = sum( - current != target for current, target in zip(self.current_placement, self.target_placement) - ) - logger.info( - "eplb started steps=%s changed_layer_count=%s changed_slot_count=%s", - self.steps, - changed_layer_count, - sum(len(transfer_batch) for transfer_batch in self.pending_transfer_batches), - ) + + pending_transfer_batches = self._transfer_planner.result + assert pending_transfer_batches is not None + self.pending_transfer_batches = pending_transfer_batches + self._transfer_planner = None + if not self.pending_transfer_batches: + raise RuntimeError("planned EPLB rearrangement must contain at least one transfer") + + self.state = EPLBManagerState.TRANSFERRING + if self.global_rank == 0: + assert self.target_placement is not None + changed_layer_count = sum( + current != target for current, target in zip(self.current_placement, self.target_placement) + ) + logger.info( + "eplb started steps=%s changed_layer_count=%s changed_slot_count=%s", + self.steps, + changed_layer_count, + sum(len(transfer_batch) for transfer_batch in self.pending_transfer_batches), + ) def _step_transferring(self) -> None: """启动或轮询一个传输批次;整批完成后再统一提交。""" - if not hasattr(self, "rebalance_started_at"): - self.rebalance_started_at = time.time() + if self.rebalance_started_at is None: + self.rebalance_started_at = time.monotonic() # 没有活动批次时,所有 rank 根据相同的 pending 列表构造下一批任务。 - if not hasattr(self, "active_transfer_batch"): + if self.active_transfer_batch is None: + assert self.pending_transfer_batches is not None transfer_batch = self.pending_transfer_batches.pop(0) if self.pending_transfer_batches else [] # 空批次表示公共任务列表已经耗尽,所有 rank 可以同时结束重排。 if not transfer_batch: + assert self.target_placement is not None self.current_placement = self.target_placement - elapsed = time.time() - self.rebalance_started_at + elapsed = time.monotonic() - self.rebalance_started_at if self.global_rank == 0: self._persist_current_placement() self._clear_prefill_route_samples() self.completed_rebalance_count += 1 - del self.pending_transfer_batches - del self.target_placement - del self.rebalance_started_at + self.pending_transfer_batches = None + self.target_placement = None + self.rebalance_started_at = None self.state = EPLBManagerState.COLLECTING if self.global_rank == 0: logger.info( @@ -468,10 +484,10 @@ def _step_transferring(self) -> None: # row,必须等整批传输完成后再统一覆盖 live 权重。 self.active_transfers = [ PinnedMemoryEPLBTransfer( - self._weights, - self.transfer_group, - self.global_rank, - transfer_info, + weights=self._weights, + transfer_group=self.transfer_group, + current_global_rank=self.global_rank, + transfer_info=transfer_info, ) for transfer_info in transfer_batch if self.global_rank in (transfer_info.source_rank, transfer_info.dest_rank) @@ -487,32 +503,32 @@ def _poll_transfer_batch(self) -> None: # 每个 rank 只负责自己参与的任务;不参与当前批次的 rank,其本地任务 # 列表为空,all([]) 自然为 True。所有 rank 汇总一个布尔值即可判断整批 # 是否完成,无需重复传输并逐条匹配 EPLBTransferInfo。 - local_finished = all(transfer.is_finished() for transfer in self.active_transfers) - finished_by_rank = [False] * self.world_size - dist.all_gather_object(finished_by_rank, local_finished, group=self.control_group) - if not all(finished_by_rank): + assert self.active_transfers is not None + if not self._all_ranks_finished(all(transfer.is_finished() for transfer in self.active_transfers)): return - else: - # 所有 rank 使用相同的批次顺序提交,因此全局 placement 和 metadata - # 始终一致;只有 destination rank 会额外写入实际专家权重。 - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - # 专家权重和路由 metadata 都由 overlap stream 上的 MoE forward - # 读取。将整批写操作排到同一条 stream,便可自然等待此前的 forward, - # 并保证后续 forward 只能看到完整提交后的权重与 metadata。 - with torch.cuda.stream(g_infer_context.get_overlap_stream()): - for transfer_info in self.active_transfer_batch: - self._commit_transfer(transfer_info) - for layer_index in {transfer_info.layer_index for transfer_info in self.active_transfer_batch}: - self._publish_layer_metadata(layer_index) - - del self.active_transfers - del self.active_transfer_batch + + # 所有 rank 使用相同的批次顺序提交,因此全局 placement 和 metadata + # 始终一致;只有 destination rank 会额外写入实际专家权重。 + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + # 专家权重和路由 metadata 都由 overlap stream 上的 MoE forward + # 读取。将整批写操作排到同一条 stream,便可自然等待此前的 forward, + # 并保证后续 forward 只能看到完整提交后的权重与 metadata。 + assert self.active_transfer_batch is not None + with torch.cuda.stream(g_infer_context.get_overlap_stream()): + for transfer_info in self.active_transfer_batch: + self._commit_transfer(transfer_info) + for layer_index in {transfer_info.layer_index for transfer_info in self.active_transfer_batch}: + self._publish_layer_metadata(layer_index) + + self.active_transfers = None + self.active_transfer_batch = None def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: """把一条已完成传输提交到 live 权重和完整布局。""" is_destination_rank = transfer_info.dest_rank == self.global_rank if is_destination_rank: + assert self.active_transfers is not None active_transfer = next( (transfer for transfer in self.active_transfers if transfer.transfer_info == transfer_info), None, @@ -534,6 +550,16 @@ def _commit_transfer(self, transfer_info: EPLBTransferInfo) -> None: transfer_info.dest_local_expert_index ] = transfer_info.source_logical_expert_id + def _all_ranks_finished(self, local_finished: bool) -> bool: + """通过控制通信组判断所有 rank 的当前后台任务是否完成。""" + finished_by_rank = [False] * self.world_size + dist.all_gather_object( + finished_by_rank, + local_finished, + group=self.control_group, + ) + return all(finished_by_rank) + def _publish_layer_metadata(self, layer_index: int) -> None: """在整批槽位更新完成后发布该层路由 metadata。""" layer_impl = self._eplb_impls[layer_index] diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index 8c4c73509b..8079353bfd 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -358,7 +358,7 @@ def test_eplb_planner_builds_legal_concrete_slot_layout(): load[:, :, :, 0] = 1000 load[:, :, :, 4] = 500 - result = planner.plan(load, current) + result = planner.plan(logical_expert_load_samples=load, current_placement=current) placement = result[0] for row in placement: @@ -374,7 +374,10 @@ def test_eplb_planner_returns_deterministic_layout_for_zero_load_experts(): planner = GreedyEPLBPlanner(2, 1) current = [[[0, 1, 3], [2, 3, 1]]] - result = planner.plan(_planner_load([[0, 0, 0, 0]], world_size=2), current) + result = planner.plan( + logical_expert_load_samples=_planner_load([[0, 0, 0, 0]], world_size=2), + current_placement=current, + ) assert result == [[[0, 1, 3], [2, 0, 1]]] @@ -393,9 +396,12 @@ def test_eplb_planner_aggregates_rank_and_sample_dimensions(): equivalent_raw_load = torch.zeros_like(raw_load) equivalent_raw_load[0, :, 0] = aggregated_load - result = planner.plan(raw_load, current) + result = planner.plan(logical_expert_load_samples=raw_load, current_placement=current) - assert result == planner.plan(equivalent_raw_load, current) + assert result == planner.plan( + logical_expert_load_samples=equivalent_raw_load, + current_placement=current, + ) def test_eplb_planner_plans_each_layer_independently_then_combines_results(): @@ -411,7 +417,7 @@ def test_eplb_planner_plans_each_layer_independently_then_combines_results(): world_size=2, ) - result = planner.plan(load, current) + result = planner.plan(logical_expert_load_samples=load, current_placement=current) assert result == [ [[0, 1, 3], [2, 0, 1]], @@ -423,7 +429,10 @@ def test_eplb_planner_iteratively_places_hot_expert_on_idle_rank(): planner = GreedyEPLBPlanner(2, 1) current = [[[0, 1, 3], [2, 3, 1]]] - result = planner.plan(_planner_load([[1000, 1, 1, 1]], world_size=2), current) + result = planner.plan( + logical_expert_load_samples=_planner_load([[1000, 1, 1, 1]], world_size=2), + current_placement=current, + ) assert result == [[[0, 1, 3], [2, 0, 1]]] @@ -433,8 +442,8 @@ def test_eplb_planner_repeatedly_splits_the_hottest_remaining_expert(): current = _initial_expert_placement(8, 4, 3).unsqueeze(0).tolist() result = planner.plan( - _planner_load([[1000, 900, 800, 700, 1, 1, 1, 1]], world_size=4), - current, + logical_expert_load_samples=_planner_load([[1000, 900, 800, 700, 1, 1, 1, 1]], world_size=4), + current_placement=current, ) replica_counts = [sum(expert in row for row in result[0]) for expert in range(8)] @@ -567,8 +576,8 @@ def test_eplb_planner_keeps_selected_experts_in_their_current_slots(): current = _initial_expert_placement(8, 4, 2).unsqueeze(0).tolist() result = planner.plan( - _planner_load([[50, 98, 54, 6, 34, 66, 63, 52]], world_size=4), - current, + logical_expert_load_samples=_planner_load([[50, 98, 54, 6, 34, 66, 63, 52]], world_size=4), + current_placement=current, ) # 只要专家仍分配在同一个 rank,就保留其原物理槽位。 @@ -592,7 +601,7 @@ def test_eplb_planner_fills_every_rank_with_distinct_nonlocal_experts(): ) load = load_by_layer_and_rank.permute(1, 0, 2).unsqueeze(dim=2) - result = planner.plan(load, current) + result = planner.plan(logical_expert_load_samples=load, current_placement=current) assert len(result) == len(current) assert all(len(actual) == len(expected) for actual, expected in zip(result[0], current[0])) @@ -628,7 +637,7 @@ def test_eplb_planner_supports_multiple_redundant_experts_per_rank(): world_size=4, ) - result = planner.plan(load, current) + result = planner.plan(logical_expert_load_samples=load, current_placement=current) for row in result[0]: assert len(row) == len(set(row)) == 7 @@ -891,9 +900,8 @@ def start(self): assert len(tasks) == 1 assert tasks[0].started assert tasks[0].kwargs["load_gather_group"] is manager.load_gather_group - assert tasks[0].kwargs["world_size"] == manager.world_size assert torch.equal( - tasks[0].kwargs["local_load"], + tasks[0].kwargs["local_load_samples"], torch.tensor([[[10, 11]], [[40, 41]]], dtype=torch.int64), ) assert torch.equal(counters[0], torch.tensor([[10, 11]], dtype=torch.int64)) @@ -917,10 +925,10 @@ def all_gather(output, local, *, group): output[1].copy_(local + 100) monkeypatch.setattr(load_gather_module.dist, "all_gather", all_gather) + monkeypatch.setattr(load_gather_module.dist, "get_world_size", lambda *, group: 2) task = load_gather_module.EPLBLoadGatherTask( - local_load=local_load, + local_load_samples=local_load, load_gather_group=load_gather_group, - world_size=2, ) task._run() @@ -940,8 +948,16 @@ def test_manager_delegates_distribution_planning_to_planner_class(): logical_load = torch.tensor([[[[10, 20]]]]) calls = [] planned_placement = [[[0, 1]]] - planner = SimpleNamespace(plan=lambda load, placement: (calls.append((load, placement)) or planned_placement)) - task = plan_module.EPLBPlanTask(planner, logical_load, current_placement) + planner = SimpleNamespace( + plan=lambda **kwargs: ( + calls.append((kwargs["logical_expert_load_samples"], kwargs["current_placement"])) or planned_placement + ) + ) + task = plan_module.EPLBPlanTask( + planner=planner, + logical_expert_load_samples=logical_load, + current_placement=current_placement, + ) task._run() @@ -953,13 +969,13 @@ def test_manager_delegates_distribution_planning_to_planner_class(): def test_plan_task_exits_process_on_failure(monkeypatch): - def fail(_load, _placement): + def fail(**_kwargs): raise RuntimeError("planning boom") task = plan_module.EPLBPlanTask( - SimpleNamespace(plan=fail), - torch.tensor([[[[10, 20]]]]), - [[[1]]], + planner=SimpleNamespace(plan=fail), + logical_expert_load_samples=torch.tensor([[[[10, 20]]]]), + current_placement=[[[1]]], ) exits = [] logs = [] @@ -981,14 +997,14 @@ def test_transfer_planner_combines_all_layer_batches(monkeypatch): ] calls = [] - def build_plan(*args): - calls.append(args) - return [[transfer_infos[args[2]]]] + def build_plan(**kwargs): + calls.append(kwargs) + return [[transfer_infos[kwargs["layer_index"]]]] monkeypatch.setattr(transfer_planner_module, "build_transfer_plan", build_plan) planner = transfer_planner_module.EPLBTransferPlanner( - current_placement, - target_placement, + current_placement=current_placement, + target_placement=target_placement, num_logical_experts=4, world_size=2, ) @@ -998,18 +1014,30 @@ def build_plan(*args): assert planner.status == "succeeded" assert planner.result == [[transfer_infos[0]], [transfer_infos[1]]] assert calls == [ - (current_placement[0], target_placement[0], 0, 4, 2), - (current_placement[1], target_placement[1], 1, 4, 2), + { + "current_placement": current_placement[0], + "target_placement": target_placement[0], + "layer_index": 0, + "num_logical_experts": 4, + "world_size": 2, + }, + { + "current_placement": current_placement[1], + "target_placement": target_placement[1], + "layer_index": 1, + "num_logical_experts": 4, + "world_size": 2, + }, ] def test_transfer_planner_exits_process_on_failure(monkeypatch): - def fail(*_args): + def fail(**_kwargs): raise RuntimeError("transfer planning boom") planner = transfer_planner_module.EPLBTransferPlanner( - [[[0], [1]]], - [[[1], [0]]], + current_placement=[[[0], [1]]], + target_placement=[[[1], [0]]], num_logical_experts=2, world_size=2, ) @@ -1115,13 +1143,13 @@ def test_publish_expert_load_metrics(): ] -def test_publish_rebalance_compute_metrics_from_global_load(): +def test_publish_rebalance_compute_metrics_from_sample_load(): calls = [] metric_client = SimpleNamespace(gauge_set=lambda name, value: calls.append((name, value))) eplb_metrics.publish_rebalance_compute_metrics( metric_client=metric_client, - global_load=torch.tensor([[384, 128, 128, 128]]), + sample_load=torch.tensor([[384, 128, 128, 128]]), current_placement=[[[0, 1, 2], [1, 2, 3]]], target_placement=[[[0, 1, 2], [0, 2, 3]]], expert_alignment=128, @@ -1221,11 +1249,11 @@ def all_gather_object(output, local_finished, **kwargs): assert seen["group"] is manager.control_group assert seen["local_finished"] is True assert manager.state is manager_module.EPLBManagerState.PLAN_PLACEMENT - assert manager._planning_load_samples is gathered_load + assert manager._pending_plan_load_samples is gathered_load # metrics 只汇总各 rank 的 sample 0;rank 1 在 sample 1 中的负载 # 不应累加进来。 - assert torch.equal(manager._planning_global_load, torch.full((1, 4), 150, dtype=torch.int64)) - assert not hasattr(manager, "_load_gather_task") + assert torch.equal(manager._metric_sample_load, torch.full((1, 4), 150, dtype=torch.int64)) + assert manager._load_gather_task is None def test_manager_wait_load_gather_does_not_advance_until_every_rank_finishes(monkeypatch): @@ -1858,6 +1886,9 @@ def is_finished(self): manager.state = manager_module.EPLBManagerState.TRANSFERRING manager.max_rebalance_count = -1 manager.completed_rebalance_count = 0 + manager.rebalance_started_at = None + manager.active_transfer_batch = None + manager.active_transfers = None committed = [] active_streams = [] cleared_route_counters = [] @@ -1895,7 +1926,7 @@ def __exit__(self, *_args): monkeypatch.setattr( manager_module, "PinnedMemoryEPLBTransfer", - lambda _weights, _group, _rank, transfer_info: Transfer(transfer_info), + lambda **kwargs: Transfer(kwargs["transfer_info"]), ) gathered_states = [ @@ -1923,7 +1954,7 @@ def all_gather_object(output, local_state, **_kwargs): manager.active_transfers[0].finished = True manager._step_transferring() assert committed == [remote_info, local_info0] - assert not hasattr(manager, "active_transfers") + assert manager.active_transfers is None manager._step_transferring() assert starts == [local_info0, local_info1] @@ -1938,9 +1969,9 @@ def all_gather_object(output, local_state, **_kwargs): [[0, 1, 2], [2, 3, 3], [4, 5, 2], [6, 7, 0]], [[0, 1, 4], [2, 3, 5], [4, 5, 6], [6, 7, 1]], ] - assert not hasattr(manager, "pending_transfer_batches") - assert not hasattr(manager, "target_placement") - assert not hasattr(manager, "rebalance_started_at") + assert manager.pending_transfer_batches is None + assert manager.target_placement is None + assert manager.rebalance_started_at is None assert cleared_route_counters == [True] assert local_states == [False, True, True] assert used_streams == [overlap_stream, overlap_stream] @@ -1958,6 +1989,9 @@ def test_manager_returns_to_collecting_after_reaching_rebalance_limit(): manager.pending_transfer_batches = [] manager.max_rebalance_count = 1 manager.completed_rebalance_count = 0 + manager.rebalance_started_at = None + manager.active_transfer_batch = None + manager.active_transfers = None manager._clear_prefill_route_samples = lambda: None persisted_placements = [] manager._persist_current_placement = lambda: persisted_placements.append(manager.current_placement) @@ -1998,8 +2032,8 @@ def test_wait_plan_finish_publishes_before_and_after_metrics(monkeypatch): manager.control_group = object() manager.current_placement = [[[0, 1], [2, 3]]] target_placement = [[[0, 2], [1, 3]]] - manager._planning_load_samples = torch.tensor([[[[4, 3, 2, 1]]]], dtype=torch.int64) - manager._planning_global_load = torch.tensor([[4, 3, 2, 1]], dtype=torch.int64) + manager._pending_plan_load_samples = None + manager._metric_sample_load = torch.tensor([[4, 3, 2, 1]], dtype=torch.int64) manager._plan_task = SimpleNamespace( is_finished=lambda: True, result=target_placement, @@ -2009,7 +2043,7 @@ def test_wait_plan_finish_publishes_before_and_after_metrics(monkeypatch): monkeypatch.setattr( eplb_metrics, "publish_rebalance_compute_metrics", - lambda **kwargs: published.append((kwargs["global_load"], kwargs["target_placement"])), + lambda **kwargs: published.append((kwargs["sample_load"], kwargs["target_placement"])), ) monkeypatch.setattr(manager_module.dist, "broadcast_object_list", lambda _values, **_kwargs: None) @@ -2020,9 +2054,9 @@ def test_wait_plan_finish_publishes_before_and_after_metrics(monkeypatch): assert len(published) == 1 assert torch.equal(published[0][0], torch.tensor([[4, 3, 2, 1]], dtype=torch.int64)) assert published[0][1] is target_placement - assert not hasattr(manager, "_planning_load_samples") - assert not hasattr(manager, "_planning_global_load") - assert not hasattr(manager, "_plan_task") + assert manager._pending_plan_load_samples is None + assert manager._metric_sample_load is None + assert manager._plan_task is None @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @@ -2170,6 +2204,7 @@ def test_manager_plans_transfers_asynchronously_before_entering_transferring(mon manager.transfer_group = object() manager.world_size = 2 manager.num_logical_experts = 4 + manager.pending_transfer_batches = None placement = [ [[0, 1, 2], [2, 3, 3]], [[0, 1, 3], [2, 3, 0]], @@ -2191,11 +2226,11 @@ def test_manager_plans_transfers_asynchronously_before_entering_transferring(mon transfer_planners = [] class TransferPlanner: - def __init__(self, current, target, num_logical_experts, world_size): - self.current = current - self.target = target - self.num_logical_experts = num_logical_experts - self.world_size = world_size + def __init__(self, **kwargs): + self.current = kwargs["current_placement"] + self.target = kwargs["target_placement"] + self.num_logical_experts = kwargs["num_logical_experts"] + self.world_size = kwargs["world_size"] self.result = [[transfer_infos[0]], [transfer_infos[1]]] self.started = False self.finished = False @@ -2211,7 +2246,7 @@ def is_finished(self): monkeypatch.setattr( manager_module, "PinnedMemoryEPLBTransfer", - lambda *_args: pytest.fail("transfer object must not be built while planning transfers"), + lambda **_kwargs: pytest.fail("transfer object must not be built while planning transfers"), ) manager._step_wait_plan_placement_finished() @@ -2219,7 +2254,7 @@ def is_finished(self): assert manager.state is manager_module.EPLBManagerState.PLAN_TRANSFER assert manager.target_placement is placement assert transfer_planners == [] - assert not hasattr(manager, "pending_transfer_batches") + assert manager.pending_transfer_batches is None manager.step() @@ -2242,16 +2277,14 @@ def gather_finished(output, local_finished, **_kwargs): manager.step() assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_TRANSFER_FINISHED - assert not hasattr(manager, "pending_transfer_batches") + assert manager.pending_transfer_batches is None remote_finished = True manager.step() assert manager.state is manager_module.EPLBManagerState.TRANSFERRING assert manager.pending_transfer_batches == [[transfer_infos[0]], [transfer_infos[1]]] - assert not hasattr(manager, "_transfer_planner") - assert not hasattr(manager, "active_transfer") - assert not hasattr(manager, "target_metadata") + assert manager._transfer_planner is None def test_manager_evaluation_with_insufficient_tokens_returns_to_collecting(monkeypatch): @@ -2306,8 +2339,8 @@ def test_manager_evaluation_with_enough_tokens_enters_planning(monkeypatch): class LoadGatherTask: def __init__(self, **kwargs): - self.local_load = kwargs["local_load"] - self.result = self.local_load.unsqueeze(0) + self.local_load_samples = kwargs["local_load_samples"] + self.result = self.local_load_samples.unsqueeze(0) self.started = False def start(self): @@ -2325,10 +2358,10 @@ def is_finished(self): ) class PlanTask: - def __init__(self, planner, logical_expert_load, current_placement): - self.planner = planner - self.logical_expert_load = logical_expert_load - self.current_placement = current_placement + def __init__(self, **kwargs): + self.planner = kwargs["planner"] + self.logical_expert_load_samples = kwargs["logical_expert_load_samples"] + self.current_placement = kwargs["current_placement"] self.started = False plan_tasks.append(self) @@ -2349,31 +2382,32 @@ def start(self): assert manager.state is manager_module.EPLBManagerState.WAIT_LOAD_GATHER_FINISHED assert manager._load_gather_task.started - assert torch.equal(manager._load_gather_task.local_load, local_load.unsqueeze(0)) + assert torch.equal(manager._load_gather_task.local_load_samples, local_load.unsqueeze(0)) assert len(published_loads) == 1 assert torch.equal(published_loads[0], local_load) assert plan_tasks == [] - assert not hasattr(manager, "_plan_task") + assert getattr(manager, "_plan_task", None) is None manager.step() assert manager.state is manager_module.EPLBManagerState.PLAN_PLACEMENT - planning_load_samples = manager._planning_load_samples + planning_load_samples = manager._pending_plan_load_samples assert torch.equal(planning_load_samples, local_load.unsqueeze(0).unsqueeze(0)) - assert torch.equal(manager._planning_global_load, local_load) - assert not hasattr(manager, "_load_gather_task") + assert torch.equal(manager._metric_sample_load, local_load) + assert manager._load_gather_task is None assert plan_tasks == [] manager.step() assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED - assert plan_tasks[0].logical_expert_load is planning_load_samples - assert torch.equal(plan_tasks[0].logical_expert_load, local_load.unsqueeze(0).unsqueeze(0)) + assert plan_tasks[0].logical_expert_load_samples is planning_load_samples + assert torch.equal(plan_tasks[0].logical_expert_load_samples, local_load.unsqueeze(0).unsqueeze(0)) assert plan_tasks[0].planner is manager.planner assert plan_tasks[0].current_placement is manager.current_placement assert plan_tasks[0].started assert manager._plan_task is plan_tasks[0] - assert torch.equal(manager._planning_global_load, local_load) + assert manager._pending_plan_load_samples is None + assert torch.equal(manager._metric_sample_load, local_load) assert len(published_loads) == 1 @@ -2417,9 +2451,9 @@ def test_nonzero_rank_enters_plan_wait_without_starting_planner(): manager.step() assert manager.state is manager_module.EPLBManagerState.WAIT_PLAN_PLACEMENT_FINISHED - assert not hasattr(manager, "_plan_task") - assert not hasattr(manager, "_planning_load_samples") - assert not hasattr(manager, "_planning_global_load") + assert getattr(manager, "_plan_task", None) is None + assert getattr(manager, "_pending_plan_load_samples", None) is None + assert getattr(manager, "_metric_sample_load", None) is None def test_manager_planning_without_changes_returns_to_collecting(monkeypatch): @@ -2446,7 +2480,7 @@ def broadcast(values, **_kwargs): assert manager.state is manager_module.EPLBManagerState.COLLECTING assert manager.next_evaluation_step == 31 - assert not hasattr(manager, "_plan_task") + assert getattr(manager, "_plan_task", None) is None def test_manager_exposes_one_lifecycle_step_entrypoint(): @@ -2709,8 +2743,8 @@ def all_gather_object(output, local_expert_ids_by_layer, group): lambda *_args, **_kwargs: pytest.fail("manager initialization must not save the placement"), ) manager = manager_module.EPLBManager(type("Model", (), {})(), config_path="/tmp/eplb.json") - assert not hasattr(manager, "_plan_task") - assert not hasattr(manager, "pending_transfer_batches") + assert manager._plan_task is None + assert manager.pending_transfer_batches is None assert manager.state is manager_module.EPLBManagerState.COLLECTING assert (manager.load_gather_group, manager.control_group, manager.transfer_group) == tuple(groups) assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 36825d7e5e..61a5448ea3 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -158,14 +158,22 @@ def _worker(rank, port): target = [[0, 1, 3], [2, 3, 1]] for expected_layer in range(2): transfer_batches = build_transfer_plan( - current, - target, - expected_layer, + current_placement=current, + target_placement=target, + layer_index=expected_layer, num_logical_experts=4, world_size=2, ) for transfer_batch in transfer_batches: - transfers = [PinnedMemoryEPLBTransfer(weights, transfer_group, rank, info) for info in transfer_batch] + transfers = [ + PinnedMemoryEPLBTransfer( + weights=weights, + transfer_group=transfer_group, + current_global_rank=rank, + transfer_info=info, + ) + for info in transfer_batch + ] for transfer in transfers: assert all(buffer.pinned_row.is_pinned() for buffer in transfer.tensor_buffers) transfer.start() @@ -190,11 +198,25 @@ def _worker(rank, port): # 主槽位互换会形成覆盖环。两个方向必须同时完成 GPU -> pinned memory # 传输后才能 commit,验证同一 rank 上并发的 send/recv 任务可以正常结束。 swap_target = [[0, 3, 2], [2, 1, 0]] - swap_plan = build_transfer_plan(current, swap_target, 0, num_logical_experts=4, world_size=2) + swap_plan = build_transfer_plan( + current_placement=current, + target_placement=swap_target, + layer_index=0, + num_logical_experts=4, + world_size=2, + ) assert len(swap_plan) == 1 swap_infos = swap_plan[0] assert len(swap_infos) == 2 - swap_transfers = [PinnedMemoryEPLBTransfer(weights, transfer_group, rank, info) for info in swap_infos] + swap_transfers = [ + PinnedMemoryEPLBTransfer( + weights=weights, + transfer_group=transfer_group, + current_global_rank=rank, + transfer_info=info, + ) + for info in swap_infos + ] for transfer in swap_transfers: transfer.start() for transfer in swap_transfers: @@ -262,7 +284,12 @@ def _many_concurrent_p2p_worker(rank, port): if rank in (transfer_info.source_rank, transfer_info.dest_rank) ] transfers = [ - PinnedMemoryEPLBTransfer(weights, transfer_group, rank, transfer_info) + PinnedMemoryEPLBTransfer( + weights=weights, + transfer_group=transfer_group, + current_global_rank=rank, + transfer_info=transfer_info, + ) for transfer_info in local_transfer_infos ] assert len(transfer_infos) == 768