From a20249d367c2aeaa5533f8d32c6e7088004f6f27 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 22 Jul 2026 15:49:26 +0800 Subject: [PATCH 1/4] 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 e80f2b552f..6d367dc390 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -73,6 +73,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"] @@ -374,9 +375,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() @@ -863,6 +869,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 a5ba656c9c..5272ba3492 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 58d4d45514..9b14f04645 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 d72e63d724..234f56b471 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -725,6 +725,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 5ccaa5a401..58237df100 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -178,6 +178,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 16aaf375aef0bdaec78d5875ff8cb23bfafe3997 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 5 Aug 2026 17:53:26 +0800 Subject: [PATCH 2/4] 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 fbb63d09f0..def46e4791 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -641,17 +641,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 69edf50a86..1f4910458d 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -644,17 +644,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 5272ba3492..21bcc58a49 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 -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 234f56b471..792961fc6e 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -730,17 +730,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 58237df100..c4dcc6b5c7 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -179,8 +179,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 4fed9509a9..b59e795e50 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 0de616b2df5fcb12880797ea9eaed0e34cc47b24 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 22 Jul 2026 15:49:26 +0800 Subject: [PATCH 3/4] 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 6d367dc390..b468a8c91b 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -73,6 +73,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"] @@ -376,13 +377,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() @@ -869,7 +872,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 21bcc58a49..ebf96b076c 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 -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 1c01cbd638..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 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_autotune_warmup(): - 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 792961fc6e..849917fec2 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -730,6 +730,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 37fe837ad1..fea5452fc4 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__) @@ -149,6 +150,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 c4dcc6b5c7..37b8a85ee5 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -179,6 +179,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 28f2abf74b..da23038e51 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -251,14 +251,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 4d09476849..448ef7b9b8 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 9a81927bc1..8e7934695e 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 b59e795e50..52866e71ad 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 2638fd88f4fd662efdaf51c17011bc24aea38f9d Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 2 Sep 2026 19:06:46 +0800 Subject: [PATCH 4/4] 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 | 35 ++++++++ 6 files changed, 106 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 fea5452fc4..45123807be 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -158,7 +158,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 da23038e51..9e1ac0e53b 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -48,6 +48,9 @@ from lightllm.server.pd_io_struct import PDChunckedTransTaskRet 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: @@ -345,7 +348,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..35d71d46b0 100644 --- a/unit_tests/server/test_api_start_eplb.py +++ b/unit_tests/server/test_api_start_eplb.py @@ -26,3 +26,38 @@ 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, "get_model_type", lambda model_dir: "llama") + 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, "register_process_tree", lambda process: None) + + api_start._launch_subprocesses(args)