diff --git a/python/sglang/jit_kernel/csrc/lora/moe_lora_align_kernel.cu b/python/sglang/jit_kernel/csrc/lora/moe_lora_align_kernel.cu index 5f6145286..3e365f801 100644 --- a/python/sglang/jit_kernel/csrc/lora/moe_lora_align_kernel.cu +++ b/python/sglang/jit_kernel/csrc/lora/moe_lora_align_kernel.cu @@ -253,7 +253,9 @@ SGL_DEVICE void _count_and_sort_expert_tokens( for (size_t i = tid; i < numel; i += stride) { int32_t expert_id = topk_ids[i]; - if (expert_id >= num_experts) { + // Under EP, StandardDispatcher writes -1 for experts not owned by this + // rank; must filter the sentinel before indexing cumsum/sorted buffers. + if (expert_id < 0 || expert_id >= num_experts) { continue; } diff --git a/python/sglang/srt/lora/mem_pool.py b/python/sglang/srt/lora/mem_pool.py index 34444edfc..bda008068 100644 --- a/python/sglang/srt/lora/mem_pool.py +++ b/python/sglang/srt/lora/mem_pool.py @@ -1,10 +1,17 @@ import logging import re -from typing import Callable, Dict, Iterable, List, Optional, Set, Tuple, Union +from typing import Callable, Dict, Iterable, Iterator, List, Optional, Set, Tuple, Union import torch -from sglang.srt.distributed import divide, get_pp_group +from sglang.srt.distributed import ( + divide, + get_moe_expert_parallel_rank, + get_moe_expert_parallel_world_size, + get_moe_tensor_parallel_rank, + get_moe_tensor_parallel_world_size, + get_pp_group, +) from sglang.srt.lora.eviction_policy import get_eviction_policy from sglang.srt.lora.layers import BaseLayerWithLoRA from sglang.srt.lora.lora import LoRAAdapter @@ -46,6 +53,43 @@ class EmptySlot: EMPTY_SLOT = EmptySlot() +def _get_moe_ep_context() -> Tuple[int, int]: + """Return `(moe_ep_size, moe_ep_rank)`, or `(1, 0)` if the MoE EP group + is not initialized (hermetic tests or pure-TP launches).""" + try: + return get_moe_expert_parallel_world_size(), get_moe_expert_parallel_rank() + except Exception: # pragma: no cover - MoE EP group not initialized + return 1, 0 + + +def _get_moe_tp_context() -> Tuple[int, int]: + """Return `(moe_tp_size, moe_tp_rank)`, or `(1, 0)` if the MoE TP group + is not initialized. Under `--tp N --ep N` the outer attention TP group + is consumed entirely by EP, leaving `moe_tp_size == 1`, so per-expert + MoE weights are NOT sharded along their inner dim even though attention + weights are.""" + try: + return get_moe_tensor_parallel_world_size(), get_moe_tensor_parallel_rank() + except Exception: # pragma: no cover - MoE TP group not initialized + return 1, 0 + + +def _moe_runner_keeps_global_expert_ids() -> bool: + """True if the active MoE runner keeps global `topk_ids` instead of + remapping to local IDs. Mirrors the predicate in `StandardDispatcher`.""" + try: + from sglang.srt.layers.moe.utils import get_moe_runner_backend + + b = get_moe_runner_backend() + return ( + b.is_flashinfer_cutlass() + or b.is_flashinfer_cutedsl() + or b.is_flashinfer_trtllm_routed() + ) + except Exception: # pragma: no cover - backend not initialized + return False + + class LoRAMemoryPool: """Class for memory pool management of lora modules""" @@ -76,6 +120,32 @@ class LoRAMemoryPool: self.experts_shared_outer_loras: bool = experts_shared_outer_loras self.strict_loading: bool = strict_loading + # Under EP with a Triton/DeepGEMM runner, `StandardDispatcher` remaps + # global `topk_ids` -> local expert IDs before the MoE kernel, so + # per-expert LoRA buffers must be sized and keyed by the local slice. + # FlashInfer CUTLASS/CuteDSL/TRTLLM-routed keep global IDs, and an + # uneven expert split (`num_experts % moe_ep_size != 0`, shouldn't + # happen in practice) is also treated as globally-keyed so we don't + # silently truncate experts. + self.moe_ep_size, self.moe_ep_rank = _get_moe_ep_context() + num_experts_global = self._get_num_experts(base_model) + self.moe_use_local_expert_ids = ( + self.moe_ep_size > 1 + and not _moe_runner_keeps_global_expert_ids() + and num_experts_global % self.moe_ep_size == 0 + ) + + # Per-expert MoE weights are sharded by `moe_tp_size`, NOT the outer + # `tp_size`: `moe_tp_size = tp_size // ep_size // dp_size`, so under + # e.g. `--tp 4 --ep 4` each rank holds full-width expert weights + # (`moe_tp_size == 1`). Sizing per-expert LoRA buffers by `tp_size` + # here would yield a 4x-narrower inner dim than the adapter weight + # (which `FusedMoEWithLoRA.slice_moe_lora_{a,b}_weights` correctly + # skip-slices when `moe_tp_size <= 1`), producing a shape-mismatch + # assert during weight load. Non-MoE modules still shard by + # `tp_size` because attention TP is unchanged. + self.moe_tp_size, self.moe_tp_rank = _get_moe_tp_context() + # Initialize eviction policy self.eviction_policy = get_eviction_policy(eviction_policy) @@ -157,6 +227,53 @@ class LoRAMemoryPool: or 1 ) + def _get_num_local_experts(self, base_model: torch.nn.Module) -> int: + """Experts owned by this rank. Equals the global count when EP is + off, the runner keeps global IDs, or the split isn't even (all + three cases fold into `moe_use_local_expert_ids == False`).""" + total = self._get_num_experts(base_model) + if not self.moe_use_local_expert_ids: + return total + return total // self.moe_ep_size + + def _global_to_local_expert_id(self, global_eid: int) -> Optional[int]: + """Map a global expert id to this rank's local id, or `None` if + the expert is not owned by this rank. Pass-through when buffers + are globally-keyed.""" + if not self.moe_use_local_expert_ids: + return global_eid + local = global_eid - self.moe_ep_rank * self._num_experts_local + return local if 0 <= local < self._num_experts_local else None + + def _iter_local_expert_weights( + self, + weights: Union[torch.Tensor, Dict[int, torch.Tensor]], + ) -> Iterator[Tuple[int, torch.Tensor]]: + """Yield `(local_expert_id, weight)` pairs for per-expert MoE LoRA + inputs, filtered/remapped to this rank's slice. Accepts either a + `{global_eid: 2D tensor}` dict or a 3D `[num_experts, *, *]` tensor.""" + if isinstance(weights, dict): + for gid, w in weights.items(): + lid = self._global_to_local_expert_id(gid) + if lid is not None: + yield lid, w + return + + if isinstance(weights, torch.Tensor) and weights.dim() == 3: + total = weights.shape[0] + if self.moe_use_local_expert_ids: + start = self.moe_ep_rank * self._num_experts_local + count = max(0, min(self._num_experts_local, total - start)) + else: + start, count = 0, total + for i in range(count): + yield i, weights[start + i] + return + + raise TypeError( + f"Expected dict or 3D torch.Tensor, got {type(weights).__name__}." + ) + def _get_standard_shape( self, module_name: str, @@ -191,16 +308,19 @@ class LoRAMemoryPool: module_name, self.base_hf_config, base_model, layer_idx ) c = get_stacked_multiply(module_name) + # MoE modules shard along `moe_tp_size`, not the outer `tp_size`. + effective_tp_size = ( + self.moe_tp_size if self.is_moe_module(module_name) else self.tp_size + ) if ( - self.tp_size > 1 + effective_tp_size > 1 and module_name in ROW_PARALLELISM_LINEAR_LORA_NAMES and module_name not in REPLICATED_LINEAR_LORA_NAMES ): - input_dim = divide(input_dim, self.tp_size) + input_dim = divide(input_dim, effective_tp_size) if self.is_moe_module(module_name): - num_experts = self._get_num_experts(base_model) - expert_dim = num_experts + expert_dim = self._get_num_local_experts(base_model) if self.experts_shared_outer_loras and module_name == "gate_up_proj_moe": expert_dim = 1 return ( @@ -247,17 +367,20 @@ class LoRAMemoryPool: _, output_dim = get_hidden_dim( module_name, self.base_hf_config, base_model, layer_idx ) + # MoE modules shard along `moe_tp_size`, not the outer `tp_size`. + effective_tp_size = ( + self.moe_tp_size if self.is_moe_module(module_name) else self.tp_size + ) if ( - self.tp_size > 1 + effective_tp_size > 1 and module_name not in ROW_PARALLELISM_LINEAR_LORA_NAMES and module_name not in REPLICATED_LINEAR_LORA_NAMES ): - output_dim = divide(output_dim, self.tp_size) + output_dim = divide(output_dim, effective_tp_size) # Check if MoE module and return appropriate shape if self.is_moe_module(module_name): - num_experts = self._get_num_experts(base_model) - expert_dim = num_experts + expert_dim = self._get_num_local_experts(base_model) if self.experts_shared_outer_loras and module_name == "down_proj_moe": expert_dim = 1 return (self.max_loras_per_batch, expert_dim, output_dim, max_lora_dim) @@ -290,6 +413,10 @@ class LoRAMemoryPool: def init_buffers(self, base_model: torch.nn.Module): device = next(base_model.parameters()).device + # Cached once so the per-expert load path doesn't re-walk the HF + # config for every adapter. + self._num_experts_local: int = self._get_num_local_experts(base_model) + def init_buffer( buffer: Dict[str, List[torch.Tensor]], target_modules: Set[str], @@ -717,50 +844,31 @@ class LoRAMemoryPool: ci * lora_rank : (ci + 1) * lora_rank, : ], ) - elif isinstance(weights, torch.Tensor) and weights.dim() == 3: - for eid in range(weights.shape[0]): - # Place each component at max_rank-spaced positions - # and zero gaps so the MoE kernel (which processes - # the full max_rank) sees correct data. - target_buffer[buffer_id, eid].zero_() - expert_weight = weights[eid] - if expert_weight is not None: - for ci in range(c): - buffer_view = target_buffer[ - buffer_id, - eid, - ci * max_r : ci * max_r + lora_rank, - :, - ] - load_lora_weight_tensor( - buffer_view, - expert_weight[ - ci * lora_rank : (ci + 1) * lora_rank, : - ], - ) - elif isinstance(weights, dict): - if weights is not None: - for expert_id, expert_weight in weights.items(): - # Place each component at max_rank-spaced positions - # and zero gaps so the MoE kernel (which processes - # the full max_rank) sees correct data. - target_buffer[buffer_id, expert_id].zero_() - if expert_weight is not None: - for ci in range(c): - buffer_view = target_buffer[ - buffer_id, - expert_id, - ci * max_r : ci * max_r + lora_rank, - :, - ] - load_lora_weight_tensor( - buffer_view, - expert_weight[ - ci * lora_rank : (ci + 1) * lora_rank, : - ], - ) - else: - target_buffer[buffer_id].zero_() + elif isinstance(weights, (torch.Tensor, dict)): + # Zero first so any local-expert slot the adapter + # doesn't fill (e.g. out-of-rank under EP) is clean; + # then load owned slots at max_rank-spaced offsets so + # the MoE kernel's [:max_r] / [max_r:2*max_r] slicing + # is correct. + target_buffer[buffer_id].zero_() + for local_eid, expert_weight in self._iter_local_expert_weights( + weights + ): + if expert_weight is None: + continue + for ci in range(c): + buffer_view = target_buffer[ + buffer_id, + local_eid, + ci * max_r : ci * max_r + lora_rank, + :, + ] + load_lora_weight_tensor( + buffer_view, + expert_weight[ + ci * lora_rank : (ci + 1) * lora_rank, : + ], + ) else: buffer_view = target_buffer[buffer_id, : lora_rank * c, :] load_lora_weight_tensor(buffer_view, weights) @@ -807,26 +915,18 @@ class LoRAMemoryPool: f"type={type(weights)}, " f"shape={weights.shape if isinstance(weights, torch.Tensor) else 'N/A'}" ) - elif isinstance(weights, torch.Tensor) and weights.dim() == 3: - for eid in range(weights.shape[0]): - buffer_view = target_buffer[buffer_id, eid, :, :lora_rank] - w = weights[eid] + elif isinstance(weights, (torch.Tensor, dict)): + # Zero out slots this rank owns but the adapter + # doesn't fill (padded-out / out-of-rank experts); + # then scale+load the ones it does. + target_buffer[buffer_id].zero_() + for local_eid, w in self._iter_local_expert_weights(weights): if w is not None: w = w * lora_adapter.scaling - load_lora_weight_tensor(buffer_view, w) - # Zero beyond loaded rank — MoE kernel reads full max_rank - target_buffer[buffer_id, eid, :, lora_rank:].zero_() - elif isinstance(weights, dict): - for expert_id, expert_weight in weights.items(): buffer_view = target_buffer[ - buffer_id, expert_id, :, :lora_rank + buffer_id, local_eid, :, :lora_rank ] - w = expert_weight - if w is not None: - w = w * lora_adapter.scaling load_lora_weight_tensor(buffer_view, w) - # Zero beyond loaded rank — MoE kernel reads full max_rank - target_buffer[buffer_id, expert_id, :, lora_rank:].zero_() else: buffer_view = target_buffer[buffer_id, :, :lora_rank] load_lora_weight_tensor(buffer_view, weights) diff --git a/test/registered/unit/lora/test_mem_pool_ep_unit.py b/test/registered/unit/lora/test_mem_pool_ep_unit.py new file mode 100644 index 000000000..b37248c30 --- /dev/null +++ b/test/registered/unit/lora/test_mem_pool_ep_unit.py @@ -0,0 +1,611 @@ +"""Unit tests for LoRAMemoryPool's MoE expert-parallel (EP) handling. + +Covers the global->local expert-id remapping and per-rank buffer sizing +introduced so that per-expert MoE LoRA buffers stay aligned with the +Triton MoE runner's local-id dispatch under `--ep > 1`. + +The tests exercise the class behavior without standing up a full server +or distributed groups: `LoRAMemoryPool` is instantiated via `__new__` +and only the fields the helpers read are populated. This keeps the +tests hermetic (CPU-only, no CUDA, no MoE EP group). + +Usage: + python -m pytest test/registered/unit/lora/test_mem_pool_ep_unit.py -v +""" + +from sglang.test.ci.ci_register import register_cuda_ci + +# CPU-only unit test; no CUDA/distributed dependencies. +register_cuda_ci(est_time=4, suite="stage-b-test-1-gpu-small") + +import types +import unittest +import unittest.mock as mock + +import torch + +from sglang.srt.lora.mem_pool import ( + LoRAMemoryPool, + _get_moe_ep_context, + _get_moe_tp_context, + _moe_runner_keeps_global_expert_ids, +) + + +def _make_pool( + *, + num_experts_global: int, + moe_ep_size: int, + moe_ep_rank: int, + moe_use_local_expert_ids: bool, +) -> LoRAMemoryPool: + """Construct a minimal LoRAMemoryPool for helper-level tests. + + Bypasses `__init__` (which requires a real base model, HF config, and + device allocations) and sets only the fields consulted by the EP + helpers under test. + """ + pool = LoRAMemoryPool.__new__(LoRAMemoryPool) + pool.moe_ep_size = moe_ep_size + pool.moe_ep_rank = moe_ep_rank + pool.moe_use_local_expert_ids = moe_use_local_expert_ids + # Helpers under test in this module don't consult moe_tp_size, but set + # defaults so accidental reads don't AttributeError. + pool.moe_tp_size = 1 + pool.moe_tp_rank = 0 + if moe_use_local_expert_ids and num_experts_global % moe_ep_size == 0: + pool._num_experts_local = num_experts_global // moe_ep_size + else: + pool._num_experts_local = num_experts_global + return pool + + +def _make_fake_base_model(num_experts: int) -> torch.nn.Module: + """Return a `torch.nn.Module` whose `.config` exposes `num_experts`. + + Used by `_get_num_experts` / `_get_num_local_experts` which walk the + HF config object. No real weights needed. + """ + model = torch.nn.Linear(4, 4, bias=False) + cfg = types.SimpleNamespace(num_experts=num_experts) + model.config = cfg + return model + + +class TestNumExpertHelpers(unittest.TestCase): + """`_get_num_experts` / `_get_num_local_experts` / buffer-dim picker.""" + + def test_num_experts_read_from_config(self): + model = _make_fake_base_model(num_experts=8) + self.assertEqual(LoRAMemoryPool._get_num_experts(model), 8) + + def test_num_local_experts_no_ep(self): + pool = _make_pool( + num_experts_global=8, + moe_ep_size=1, + moe_ep_rank=0, + moe_use_local_expert_ids=False, + ) + model = _make_fake_base_model(num_experts=8) + self.assertEqual(pool._get_num_local_experts(model), 8) + + def test_num_local_experts_with_ep(self): + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=2, + moe_use_local_expert_ids=True, + ) + model = _make_fake_base_model(num_experts=8) + self.assertEqual(pool._get_num_local_experts(model), 2) + + def test_num_local_experts_with_ep_but_backend_keeps_global_ids(self): + """FlashInfer-style backends keep global topk_ids, so even under EP + the LoRA buffers must remain globally-keyed. + """ + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=2, + moe_use_local_expert_ids=False, + ) + model = _make_fake_base_model(num_experts=8) + self.assertEqual(pool._get_num_local_experts(model), 8) + + def test_uneven_split_disables_local_mapping(self): + """Shouldn't happen in practice (base MoE requires even split), but + `__init__` must fold uneven splits into `moe_use_local_expert_ids == + False` so `_get_num_local_experts` returns the global count and no + remapping happens anywhere downstream. + """ + # Simulate what `LoRAMemoryPool.__init__` would set for an uneven + # split: the divisibility guard there forces the flag to False. + pool = _make_pool( + num_experts_global=7, + moe_ep_size=4, + moe_ep_rank=0, + moe_use_local_expert_ids=False, + ) + model = _make_fake_base_model(num_experts=7) + self.assertEqual(pool._get_num_local_experts(model), 7) + + +class TestGlobalToLocalExpertId(unittest.TestCase): + """`_global_to_local_expert_id` — the per-rank filter + remap.""" + + def test_passthrough_without_ep(self): + pool = _make_pool( + num_experts_global=8, + moe_ep_size=1, + moe_ep_rank=0, + moe_use_local_expert_ids=False, + ) + for gid in range(8): + self.assertEqual(pool._global_to_local_expert_id(gid), gid) + + def test_rank0_of_ep4_owns_first_quarter(self): + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=0, + moe_use_local_expert_ids=True, + ) + # Owned: 0, 1 -> local 0, 1 + self.assertEqual(pool._global_to_local_expert_id(0), 0) + self.assertEqual(pool._global_to_local_expert_id(1), 1) + # Not owned by rank 0. + for gid in (2, 3, 4, 5, 6, 7): + self.assertIsNone(pool._global_to_local_expert_id(gid)) + + def test_rank2_of_ep4_owns_third_quarter(self): + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=2, + moe_use_local_expert_ids=True, + ) + # Owned globals 4, 5 -> local 0, 1 + self.assertEqual(pool._global_to_local_expert_id(4), 0) + self.assertEqual(pool._global_to_local_expert_id(5), 1) + for gid in (0, 1, 2, 3, 6, 7): + self.assertIsNone(pool._global_to_local_expert_id(gid)) + + def test_last_rank_owns_last_slice(self): + pool = _make_pool( + num_experts_global=128, + moe_ep_size=4, + moe_ep_rank=3, + moe_use_local_expert_ids=True, + ) + # Local 0 <-> global 96, local 31 <-> global 127. + self.assertEqual(pool._global_to_local_expert_id(96), 0) + self.assertEqual(pool._global_to_local_expert_id(127), 31) + self.assertIsNone(pool._global_to_local_expert_id(95)) + self.assertIsNone(pool._global_to_local_expert_id(128)) + + +class TestIterLocalExpertWeightsDict(unittest.TestCase): + """`_iter_local_expert_weights` with dict input (the common case).""" + + def test_passthrough_without_ep(self): + pool = _make_pool( + num_experts_global=4, + moe_ep_size=1, + moe_ep_rank=0, + moe_use_local_expert_ids=False, + ) + weights = {gid: torch.full((2,), float(gid)) for gid in range(4)} + got = {lid: w.tolist() for lid, w in pool._iter_local_expert_weights(weights)} + self.assertEqual( + got, {0: [0.0, 0.0], 1: [1.0, 1.0], 2: [2.0, 2.0], 3: [3.0, 3.0]} + ) + + def test_rank0_of_ep4_filters_and_remaps(self): + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=0, + moe_use_local_expert_ids=True, + ) + weights = {gid: torch.full((2,), float(gid)) for gid in range(8)} + got = {lid: w.tolist() for lid, w in pool._iter_local_expert_weights(weights)} + # Rank 0 sees globals 0,1 remapped to locals 0,1. + self.assertEqual(got, {0: [0.0, 0.0], 1: [1.0, 1.0]}) + + def test_rank3_of_ep4_filters_and_remaps(self): + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=3, + moe_use_local_expert_ids=True, + ) + weights = {gid: torch.full((2,), float(gid)) for gid in range(8)} + got = {lid: w.tolist() for lid, w in pool._iter_local_expert_weights(weights)} + # Rank 3 sees globals 6,7 remapped to locals 0,1. + self.assertEqual(got, {0: [6.0, 6.0], 1: [7.0, 7.0]}) + + def test_sparse_dict_only_yields_owned_experts(self): + """Adapters may only target a subset of experts. The iterator must + still correctly filter and remap whatever subset is provided. + """ + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=2, + moe_use_local_expert_ids=True, + ) + # Only globals 1, 4, 5, 7 present in adapter. + weights = { + 1: torch.full((2,), 1.0), + 4: torch.full((2,), 4.0), + 5: torch.full((2,), 5.0), + 7: torch.full((2,), 7.0), + } + # Rank 2 owns globals 4, 5 -> locals 0, 1. + got = {lid: w.tolist() for lid, w in pool._iter_local_expert_weights(weights)} + self.assertEqual(got, {0: [4.0, 4.0], 1: [5.0, 5.0]}) + + def test_no_experts_owned_yields_nothing(self): + """Rank with no matching experts in a sparse dict yields nothing, + leaves buffer zeroed. + """ + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=0, + moe_use_local_expert_ids=True, + ) + # Only globals 4, 5 present (owned by rank 2). + weights = {4: torch.full((2,), 4.0), 5: torch.full((2,), 5.0)} + got = list(pool._iter_local_expert_weights(weights)) + self.assertEqual(got, []) + + +class TestIterLocalExpertWeightsTensor(unittest.TestCase): + """`_iter_local_expert_weights` with 3D tensor input (shared-outer and + packed MoE-LoRA formats).""" + + def test_passthrough_without_ep(self): + pool = _make_pool( + num_experts_global=4, + moe_ep_size=1, + moe_ep_rank=0, + moe_use_local_expert_ids=False, + ) + # [num_experts, rank, hidden] with values carrying the expert id. + weights = torch.arange(4 * 2 * 3, dtype=torch.float32).reshape(4, 2, 3) + got = [(lid, w.clone()) for lid, w in pool._iter_local_expert_weights(weights)] + self.assertEqual([lid for lid, _ in got], [0, 1, 2, 3]) + for lid, w in got: + self.assertTrue(torch.equal(w, weights[lid])) + + def test_rank1_of_ep2_sees_upper_half(self): + pool = _make_pool( + num_experts_global=4, + moe_ep_size=2, + moe_ep_rank=1, + moe_use_local_expert_ids=True, + ) + weights = torch.arange(4 * 2 * 3, dtype=torch.float32).reshape(4, 2, 3) + got = [(lid, w.clone()) for lid, w in pool._iter_local_expert_weights(weights)] + # Rank 1 of EP=2 with 4 experts owns globals 2, 3 -> locals 0, 1. + self.assertEqual([lid for lid, _ in got], [0, 1]) + self.assertTrue(torch.equal(got[0][1], weights[2])) + self.assertTrue(torch.equal(got[1][1], weights[3])) + + def test_rank_with_partial_tensor_coverage(self): + """Defensive: tensor has fewer experts than the expected local slice + (e.g. sparse adapter). + """ + pool = _make_pool( + num_experts_global=8, + moe_ep_size=4, + moe_ep_rank=3, + moe_use_local_expert_ids=True, + ) + # Only 6 experts present in the tensor; rank 3 expects global 6,7. + # So it should still yield local 0 mapped to global 6; global 7 is + # beyond the tensor length and must be skipped safely. + weights = torch.arange(6 * 2, dtype=torch.float32).reshape(6, 2) + # Note: this is 2D, not 3D -> should raise (sanity check). + with self.assertRaises(TypeError): + list(pool._iter_local_expert_weights(weights)) + + +class TestModuleLevelHelpers(unittest.TestCase): + """`_get_moe_ep_context` / `_moe_runner_keeps_global_expert_ids` + must degrade gracefully when the MoE EP group or runner backend is + not yet initialized (e.g. in pure-TP launches or hermetic tests).""" + + def test_ep_context_defaults_when_group_uninitialized(self): + # Real process here: the MoE EP group isn't set up in a unit test. + # The helper must return (1, 0) rather than raising. + ep_size, ep_rank = _get_moe_ep_context() + self.assertEqual(ep_size, 1) + self.assertEqual(ep_rank, 0) + + def test_tp_context_defaults_when_group_uninitialized(self): + # Mirror of `_get_moe_ep_context` for the MoE TP group: if it isn't + # initialized (hermetic tests, pure-TP launches), fall back to (1, 0). + tp_size, tp_rank = _get_moe_tp_context() + self.assertEqual(tp_size, 1) + self.assertEqual(tp_rank, 0) + + def test_keeps_global_expert_ids_defaults_to_false(self): + # Without a specific flashinfer backend selected, default is False. + self.assertFalse(_moe_runner_keeps_global_expert_ids()) + + +class TestPoolInitPicksUpEpContext(unittest.TestCase): + """`LoRAMemoryPool.__init__` should read EP context from the module- + level helpers and set `moe_use_local_expert_ids` correctly.""" + + def _new_pool_with_ep( + self, + ep_size: int, + ep_rank: int, + keeps_global: bool, + num_experts: int = 8, + moe_tp_size: int = 1, + moe_tp_rank: int = 0, + tp_size: int = 1, + tp_rank: int = 0, + ) -> LoRAMemoryPool: + """Construct a pool with `__init__` called, but stop before + `init_buffers` — we only care about the EP-context state. + """ + with ( + mock.patch( + "sglang.srt.lora.mem_pool._get_moe_ep_context", + return_value=(ep_size, ep_rank), + ), + mock.patch( + "sglang.srt.lora.mem_pool._get_moe_tp_context", + return_value=(moe_tp_size, moe_tp_rank), + ), + mock.patch( + "sglang.srt.lora.mem_pool._moe_runner_keeps_global_expert_ids", + return_value=keeps_global, + ), + mock.patch.object(LoRAMemoryPool, "init_buffers", lambda self, _m: None), + ): + hf_cfg = types.SimpleNamespace( + num_hidden_layers=1, + hidden_size=8, + vocab_size=32, + num_experts=num_experts, + ) + base_model = torch.nn.Linear(8, 8, bias=False) + base_model.config = hf_cfg + return LoRAMemoryPool( + base_hf_config=hf_cfg, + max_loras_per_batch=1, + dtype=torch.bfloat16, + tp_size=tp_size, + tp_rank=tp_rank, + max_lora_rank=8, + target_modules={"qkv_proj"}, + base_model=base_model, + eviction_policy="lru", + lora_added_tokens_size=0, + ) + + def test_no_ep(self): + pool = self._new_pool_with_ep(ep_size=1, ep_rank=0, keeps_global=False) + self.assertEqual(pool.moe_ep_size, 1) + self.assertEqual(pool.moe_ep_rank, 0) + self.assertFalse(pool.moe_use_local_expert_ids) + + def test_ep4_triton_backend(self): + pool = self._new_pool_with_ep(ep_size=4, ep_rank=2, keeps_global=False) + self.assertEqual(pool.moe_ep_size, 4) + self.assertEqual(pool.moe_ep_rank, 2) + self.assertTrue(pool.moe_use_local_expert_ids) + + def test_ep4_flashinfer_cutlass_keeps_global(self): + """FlashInfer CUTLASS keeps global topk_ids, so LoRA buffers stay + globally-keyed even under EP. + """ + pool = self._new_pool_with_ep(ep_size=4, ep_rank=2, keeps_global=True) + self.assertEqual(pool.moe_ep_size, 4) + self.assertEqual(pool.moe_ep_rank, 2) + self.assertFalse(pool.moe_use_local_expert_ids) + + def test_ep_with_uneven_split_falls_back_to_global_ids(self): + """If `num_experts % ep_size != 0` (shouldn't happen in practice, + base MoE requires even split) `__init__` must fall back to + globally-keyed buffers rather than silently truncating the local + slice — otherwise non-zero ranks drop every LoRA weight. + """ + pool = self._new_pool_with_ep( + ep_size=4, ep_rank=1, keeps_global=False, num_experts=7 + ) + self.assertEqual(pool.moe_ep_size, 4) + self.assertEqual(pool.moe_ep_rank, 1) + self.assertFalse(pool.moe_use_local_expert_ids) + + def test_init_captures_moe_tp_context(self): + """`__init__` must capture moe_tp_size/rank so per-expert MoE LoRA + buffers can be sharded by the MoE-TP group (not the outer attn TP). + Under `--tp N --ep N` the MoE TP group degenerates to size 1. + """ + pool = self._new_pool_with_ep( + ep_size=4, + ep_rank=0, + keeps_global=False, + tp_size=4, + tp_rank=0, + moe_tp_size=1, + moe_tp_rank=0, + ) + self.assertEqual(pool.tp_size, 4) + self.assertEqual(pool.moe_tp_size, 1) + self.assertEqual(pool.moe_tp_rank, 0) + + +def _fake_base_model_with_hidden_dim(num_experts: int) -> torch.nn.Module: + """Fake base model that implements `get_hidden_dim` for MoE + attention + modules. Matches the signatures `LoRAMemoryPool.get_lora_{A,B}_shape` + call through `sglang.srt.lora.utils.get_hidden_dim`. + """ + + class _Model(torch.nn.Module): + def __init__(self): + super().__init__() + self.lin = torch.nn.Linear(4, 4, bias=False) + self.config = types.SimpleNamespace( + num_hidden_layers=1, + hidden_size=64, + num_attention_heads=8, + num_key_value_heads=8, + head_dim=8, + intermediate_size=256, + moe_intermediate_size=192, + vocab_size=32, + num_experts=num_experts, + ) + + def get_hidden_dim(self, module_name: str, layer_idx: int): + cfg = self.config + if module_name == "qkv_proj": + head = cfg.head_dim + return cfg.hidden_size, head * ( + cfg.num_attention_heads + cfg.num_key_value_heads * 2 + ) + if module_name == "o_proj": + return cfg.head_dim * cfg.num_attention_heads, cfg.hidden_size + if module_name == "gate_up_proj_moe": + return cfg.hidden_size, cfg.moe_intermediate_size * 2 + if module_name == "down_proj_moe": + return cfg.moe_intermediate_size, cfg.hidden_size + raise NotImplementedError(module_name) + + return _Model() + + +class TestMoeBufferShardsByMoeTp(unittest.TestCase): + """Regression: per-expert MoE LoRA buffers must shard by `moe_tp_size`, + not the outer attention `tp_size`. + + Under `--tp N --ep N` (e.g. tp=4, ep=4) `moe_tp_size == 1`, so per- + expert weights span the full MoE intermediate dim on every rank; the + corresponding LoRA buffer must match. Before the fix, the buffer was + divided by `tp_size` (= 4) while `FusedMoEWithLoRA.slice_moe_lora_*` + kept the weight full-width, producing a 4x shape-mismatch assert at + load time. Non-MoE modules still shard by the outer `tp_size`. + """ + + def _pool( + self, + *, + tp_size: int, + moe_tp_size: int, + num_experts: int = 128, + ep_size: int = 1, + ep_rank: int = 0, + ) -> LoRAMemoryPool: + pool = LoRAMemoryPool.__new__(LoRAMemoryPool) + pool.max_loras_per_batch = 2 + pool.tp_size = tp_size + pool.tp_rank = 0 + pool.moe_ep_size = ep_size + pool.moe_ep_rank = ep_rank + pool.moe_tp_size = moe_tp_size + pool.moe_tp_rank = 0 + pool.moe_use_local_expert_ids = ep_size > 1 + pool._num_experts_local = ( + num_experts // ep_size if pool.moe_use_local_expert_ids else num_experts + ) + pool.experts_shared_outer_loras = False + pool.base_hf_config = types.SimpleNamespace( + hidden_size=64, + num_attention_heads=8, + num_key_value_heads=8, + head_dim=8, + intermediate_size=256, + moe_intermediate_size=192, + ) + return pool + + def test_moe_down_proj_uses_moe_tp_not_attn_tp(self): + """down_proj_moe is row-parallel: LoRA-A input_dim = moe_inter must + be divided by `moe_tp_size`, NOT `tp_size`. This is the exact shape + that failed at load time on `--tp 4 --ep 4` before the fix. + """ + pool = self._pool(tp_size=4, moe_tp_size=1, num_experts=128, ep_size=4) + model = _fake_base_model_with_hidden_dim(num_experts=128) + num_local = 128 // 4 # 32 + # A: input_dim = moe_inter / moe_tp_size = 192 / 1 = 192 (pre-fix: 48). + self.assertEqual( + pool.get_lora_A_shape("down_proj_moe", model, 8, 0), + (2, num_local, 8, 192), + ) + # B: output_dim = hidden_size, not row-parallel -> unsharded. + self.assertEqual( + pool.get_lora_B_shape("down_proj_moe", model, 8, 0), + (2, num_local, 64, 8), + ) + + def test_moe_gate_up_proj_uses_moe_tp_not_attn_tp(self): + """gate_up_proj_moe is column-parallel: LoRA-B output_dim = + moe_inter*2 must be divided by `moe_tp_size`, not `tp_size`. + """ + pool = self._pool(tp_size=4, moe_tp_size=1, num_experts=128, ep_size=4) + model = _fake_base_model_with_hidden_dim(num_experts=128) + num_local = 128 // 4 + # A: input_dim = hidden_size, not row-parallel -> unsharded. Rank + # dim is `max_lora_dim * stacked_multiply` (2 for gate_up). + self.assertEqual( + pool.get_lora_A_shape("gate_up_proj_moe", model, 8, 0), + (2, num_local, 16, 64), + ) + # B: output_dim = moe_inter*2 / moe_tp_size = 384 / 1 = 384 (pre-fix: 96). + self.assertEqual( + pool.get_lora_B_shape("gate_up_proj_moe", model, 8, 0), + (2, num_local, 384, 8), + ) + + def test_moe_tp_gt1_still_shards_moe_dims(self): + """Under `--tp 8 --ep 4` the MoE TP group has size 2, so per-expert + weights ARE sharded along the MoE inner dim — the LoRA buffer must + follow. + """ + pool = self._pool(tp_size=8, moe_tp_size=2, num_experts=128, ep_size=4) + model = _fake_base_model_with_hidden_dim(num_experts=128) + num_local = 128 // 4 + # 192 / 2 = 96 + self.assertEqual( + pool.get_lora_A_shape("down_proj_moe", model, 8, 0), + (2, num_local, 8, 96), + ) + # 384 / 2 = 192 (B: moe_inter*2 / moe_tp_size). + self.assertEqual( + pool.get_lora_B_shape("gate_up_proj_moe", model, 8, 0), + (2, num_local, 192, 8), + ) + # A: input_dim = hidden_size, unaffected by MoE TP. + self.assertEqual( + pool.get_lora_A_shape("gate_up_proj_moe", model, 8, 0), + (2, num_local, 16, 64), + ) + + def test_non_moe_modules_unaffected_by_moe_tp(self): + """Non-MoE modules must continue to shard by the outer `tp_size`; + the MoE-TP substitution applies only to `*_moe` modules. + """ + pool = self._pool(tp_size=4, moe_tp_size=1, num_experts=128, ep_size=4) + model = _fake_base_model_with_hidden_dim(num_experts=128) + # o_proj is row-parallel: A input_dim sharded by tp_size, B unsharded. + o_a = pool.get_lora_A_shape("o_proj", model, 8, 0) + o_b = pool.get_lora_B_shape("o_proj", model, 8, 0) + # head_dim*num_heads / tp_size = 64 / 4 = 16; B output = hidden_size = 64. + self.assertEqual(o_a, (2, 8, 16)) + self.assertEqual(o_b, (2, 64, 8)) + # qkv_proj is column-parallel: A unsharded, B sharded by tp_size. + q_b = pool.get_lora_B_shape("qkv_proj", model, 8, 0) + # head_dim * (heads + 2*kv_heads) / tp_size = 8 * 24 / 4 = 48. + self.assertEqual(q_b, (2, 48, 8)) + + +if __name__ == "__main__": + unittest.main()