From f81b6df3a3fbff7203e604e935091e26e29662b1 Mon Sep 17 00:00:00 2001 From: Kurt Shuster Date: Sun, 12 Apr 2026 19:25:02 -0400 Subject: [PATCH] [lora] Fix partial MoE rank loading, VL lm_head, strict loading, deepseek on-demand (#21864) Co-authored-by: Yusheng Su --- python/sglang/srt/lora/lora_manager.py | 35 +++-- python/sglang/srt/lora/mem_pool.py | 204 +++++++++++++++++++------ python/sglang/srt/server_args.py | 12 +- 3 files changed, 186 insertions(+), 65 deletions(-) diff --git a/python/sglang/srt/lora/lora_manager.py b/python/sglang/srt/lora/lora_manager.py index 06f3b09dd..9f3ef601e 100644 --- a/python/sglang/srt/lora/lora_manager.py +++ b/python/sglang/srt/lora/lora_manager.py @@ -35,6 +35,7 @@ from sglang.srt.lora.lora_config import LoRAConfig from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.lora.mem_pool import LoRAMemoryPool from sglang.srt.lora.utils import ( + EMBEDDING_NAMES, LoRAType, auto_detect_lora_target_modules, get_normalized_target_modules, @@ -85,6 +86,9 @@ class LoRAManager: self._experts_shared_outer_override: Optional[bool] = ( server_args.experts_shared_outer_loras ) + self.lora_strict_loading: bool = getattr( + server_args, "lora_strict_loading", False + ) # LoRA backend for running sgemm kernels logger.info(f"Using {lora_backend} as backend of LoRA kernels.") @@ -507,9 +511,19 @@ class LoRAManager: ): """Infer LoRA target modules and max_lora_rank from loaded adapters if not provided.""" - self.target_modules = ( - get_normalized_target_modules(target_modules) if target_modules else set() - ) + if target_modules and target_modules == {"all"}: + self.target_modules = auto_detect_lora_target_modules(self.base_model) + self.target_modules.update(EMBEDDING_NAMES) + logger.info( + "CLI --lora-target-modules='all' resolved to %s " + "by inspecting the base model.", + sorted(self.target_modules), + ) + target_modules = self.target_modules + elif target_modules: + self.target_modules = get_normalized_target_modules(target_modules) + else: + self.target_modules = set() for lora_id, config in self.configs.items(): # Handle PEFT shorthand strings like "all-linear" or "all". @@ -682,6 +696,7 @@ class LoRAManager: eviction_policy=self.eviction_policy, lora_added_tokens_size=self.lora_added_tokens_size, experts_shared_outer_loras=self.experts_shared_outer_loras, + strict_loading=self.lora_strict_loading, ) # Initializing memory pool with base model @@ -737,17 +752,9 @@ class LoRAManager: self.base_model.lm_head = untied_lm_head for module_name, module in self.base_model.named_modules(): - # TODO (lifuhuang): in the future, we should consider generalizing the - # should_apply_lora function to support mapping by full module name instead - # of just the last part (e.g., "qkv_proj") to support scenarios with multiple - # attention stacks (e.g., multimodal models). - # See: https://github.com/sgl-project/sglang/issues/6608 - if getattr( - self.base_model, "should_apply_lora", None - ) and not self.base_model.should_apply_lora(module_name): - continue - - # Check if module should be wrapped with LoRA + # Handle embed_tokens and lm_head before the should_apply_lora gate, + # since VL models' should_apply_lora patterns only match language + # model layers and would incorrectly skip these. # Handle embed_tokens if "embed_tokens" in module_name and "embed_tokens" in self.target_modules: if isinstance(module, VocabParallelEmbedding) and not isinstance( diff --git a/python/sglang/srt/lora/mem_pool.py b/python/sglang/srt/lora/mem_pool.py index 6c923afc5..34444edfc 100644 --- a/python/sglang/srt/lora/mem_pool.py +++ b/python/sglang/srt/lora/mem_pool.py @@ -4,7 +4,7 @@ from typing import Callable, Dict, Iterable, List, Optional, Set, Tuple, Union import torch -from sglang.srt.distributed import divide +from sglang.srt.distributed import divide, 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 @@ -62,6 +62,7 @@ class LoRAMemoryPool: eviction_policy: str, lora_added_tokens_size: int, experts_shared_outer_loras: bool = False, + strict_loading: bool = False, ): self.base_hf_config: AutoConfig = base_hf_config self.num_layer: int = base_hf_config.num_hidden_layers @@ -73,6 +74,7 @@ class LoRAMemoryPool: self.max_lora_rank: int = max_lora_rank self.target_modules: Set[str] = target_modules self.experts_shared_outer_loras: bool = experts_shared_outer_loras + self.strict_loading: bool = strict_loading # Initialize eviction policy self.eviction_policy = get_eviction_policy(eviction_policy) @@ -533,6 +535,43 @@ class LoRAMemoryPool: assert lora_adapter is not None lora_rank = lora_adapter.config.r + + # Pre-validate weight names against target modules across all layers + # and embedding weights. This catches mismatches before any GPU + # buffers are mutated. + skipped_weight_names: set = set() + matched_modules: set = set() + all_weight_names: list = [] + for layer in lora_adapter.layers: + all_weight_names.extend(layer.weights.keys()) + if lora_adapter.embedding_layers: + all_weight_names.extend(lora_adapter.embedding_layers.keys()) + for name in all_weight_names: + try: + target_module = get_target_module_name(name, self.target_modules) + matched_modules.add(target_module) + except ValueError: + skipped_weight_names.add(name) + if matched_modules: + logger.info( + "LoRA adapter '%s': loaded weights for target modules %s.", + uid, + sorted(matched_modules), + ) + if skipped_weight_names: + msg = ( + f"LoRA adapter '{uid}': {len(skipped_weight_names)} weight(s) " + f"skipped because they did not match any target module in " + f"{sorted(self.target_modules)}. Skipped weights: " + f"{sorted(skipped_weight_names)}. This likely indicates a " + f"mismatch between the adapter's target modules and the base " + f"model architecture." + ) + if self.strict_loading: + raise ValueError(msg) + else: + logger.warning(msg) + for layer_id in range(self.num_layer): layer_weights = lora_adapter.layers[layer_id].weights # - Standard: module_name -> torch.Tensor @@ -576,59 +615,58 @@ class LoRAMemoryPool: else: temp_B_buffer[target_module] = weights - if self.tp_size > 1: - cur_layer_modules = lora_modules[layer_id] - for module_name, module in cur_layer_modules.items(): - # TODO (Jonahcb): check if the code can be refactored to avoid the special handling for FusedMoEWithLoRA - # Handle FusedMoEWithLoRA specially - it contains multiple target modules - from sglang.srt.lora.layers import FusedMoEWithLoRA + cur_layer_modules = lora_modules[layer_id] + for module_name, module in cur_layer_modules.items(): + # TODO (Jonahcb): check if the code can be refactored to avoid the special handling for FusedMoEWithLoRA + # Handle FusedMoEWithLoRA specially - it contains multiple target modules + from sglang.srt.lora.layers import FusedMoEWithLoRA - if isinstance(module, FusedMoEWithLoRA): - moe_target_modules = ["gate_up_proj_moe", "down_proj_moe"] - for target_module in moe_target_modules: - if temp_A_buffer.get(target_module) is not None: - temp_A_buffer[target_module] = ( - module.slice_moe_lora_a_weights( - temp_A_buffer[target_module], - self.tp_rank, - target_module, - ) + if isinstance(module, FusedMoEWithLoRA): + moe_target_modules = ["gate_up_proj_moe", "down_proj_moe"] + for target_module in moe_target_modules: + if temp_A_buffer.get(target_module) is not None: + temp_A_buffer[target_module] = ( + module.slice_moe_lora_a_weights( + temp_A_buffer[target_module], + self.tp_rank, + target_module, ) - if temp_B_buffer.get(target_module) is not None: - temp_B_buffer[target_module] = ( - module.slice_moe_lora_b_weights( - temp_B_buffer[target_module], - self.tp_rank, - target_module, - ) + ) + if temp_B_buffer.get(target_module) is not None: + temp_B_buffer[target_module] = ( + module.slice_moe_lora_b_weights( + temp_B_buffer[target_module], + self.tp_rank, + target_module, ) + ) - continue + continue - # Handle regular modules - target_module = get_target_module_name( - module_name, self.target_modules - ) + # Handle regular modules + target_module = get_target_module_name(module_name, self.target_modules) - if temp_A_buffer[target_module] is None: - # Skip weight slicing if the weight is not present in the adapter - continue + if temp_A_buffer[target_module] is None: + # Skip weight slicing if the weight is not present in the adapter + continue - # Handle standard modules - temp_A_buffer[target_module] = module.slice_lora_a_weights( - temp_A_buffer[target_module], self.tp_rank - ) - temp_B_buffer[target_module] = module.slice_lora_b_weights( - temp_B_buffer[target_module], self.tp_rank - ) + # Handle standard modules + temp_A_buffer[target_module] = module.slice_lora_a_weights( + temp_A_buffer[target_module], self.tp_rank + ) + temp_B_buffer[target_module] = module.slice_lora_b_weights( + temp_B_buffer[target_module], self.tp_rank + ) for name, weights in temp_A_buffer.items(): c = get_stacked_multiply(name) + max_r = self.max_lora_rank target_buffer = self.A_buffer[name][layer_id] if name in ["gate_up_proj_moe", "down_proj_moe"]: if self.experts_shared_outer_loras and name == "gate_up_proj_moe": if weights is None: + representative_weight = None buffer_view = target_buffer[ buffer_id, 0, : lora_rank * c, : ] @@ -640,6 +678,7 @@ class LoRAMemoryPool: f"gate_up_proj_moe lora_A has expert_dim=" f"{weights.shape[0]} (expected 1)." ) + representative_weight = weights[0] buffer_view = target_buffer[ buffer_id, 0, : lora_rank * c, : ] @@ -652,6 +691,7 @@ class LoRAMemoryPool: f"{len(weights)} entries (expected 1)." ) rep = next(iter(weights.values())) + representative_weight = rep buffer_view = target_buffer[ buffer_id, 0, : lora_rank * c, : ] @@ -662,18 +702,65 @@ class LoRAMemoryPool: f"type={type(weights)}, " f"shape={weights.shape if isinstance(weights, torch.Tensor) else 'N/A'}" ) + # Place each stacked component at max_rank-spaced + # positions so the kernel's [:max_r] / [max_r:2*max_r] + # slicing is correct. + target_buffer[buffer_id, 0].zero_() + if representative_weight is not None: + for ci in range(c): + buffer_view = target_buffer[ + buffer_id, 0, ci * max_r : ci * max_r + lora_rank, : + ] + load_lora_weight_tensor( + buffer_view, + representative_weight[ + ci * lora_rank : (ci + 1) * lora_rank, : + ], + ) 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 * c, : - ] - load_lora_weight_tensor(buffer_view, weights[eid]) + # 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): - for expert_id, expert_weight in weights.items(): - buffer_view = target_buffer[ - buffer_id, expert_id, : lora_rank * c, : - ] - load_lora_weight_tensor(buffer_view, expert_weight) + 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_() else: buffer_view = target_buffer[buffer_id, : lora_rank * c, :] load_lora_weight_tensor(buffer_view, weights) @@ -698,6 +785,8 @@ class LoRAMemoryPool: 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, 0, :, lora_rank:].zero_() elif isinstance(weights, dict) and len(weights) > 0: if len(weights) != 1: raise ValueError( @@ -710,6 +799,8 @@ class LoRAMemoryPool: if rep is not None: rep = rep * lora_adapter.scaling load_lora_weight_tensor(buffer_view, rep) + # Zero beyond loaded rank — MoE kernel reads full max_rank + target_buffer[buffer_id, 0, :, lora_rank:].zero_() else: raise ValueError( f"Unexpected weight format for shared outer down_proj_moe lora_B: " @@ -723,6 +814,8 @@ class LoRAMemoryPool: 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[ @@ -732,6 +825,8 @@ class LoRAMemoryPool: 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) @@ -779,6 +874,7 @@ class LoRAMemoryPool: elif ( target_module == "lm_head" + and lora_lm_head_module is not None and "lm_head" in name and ("lora_embedding_A" in name or "lora_A" in name) ): @@ -791,12 +887,14 @@ class LoRAMemoryPool: load_lora_weight_tensor(buffer_view, weights) elif ( target_module == "lm_head" + and lora_lm_head_module is not None and "lm_head" in name and ("lora_embedding_B" in name or "lora_B" in name) ): + assert lora_lm_head_module is not None lora_b_weights = weights # Slice B along vocab dimension for this TP rank - if self.tp_size > 1 and lora_lm_head_module is not None: + if self.tp_size > 1: lora_b_weights = lora_lm_head_module.slice_lora_b_weights( lora_b_weights, self.tp_rank ) @@ -807,6 +905,14 @@ class LoRAMemoryPool: :lora_rank, ] load_lora_weight_tensor(buffer_view, lora_b_weights) + elif target_module == "lm_head" and "lm_head" in name: + # Non-last PP stages do not own lm_head, so adapters can + # legitimately contain lm_head LoRA weights with no local + # module to load them into, otherwise we should have been able to load this weight. + assert ( + not get_pp_group().is_last_rank + ), f"Failed to load lm_head LoRA weight: {name}, this is only expected to happen on non-last PP stages." + continue else: # Zero out embedding/lm_head buffers for adapters without embedding LoRA # to avoid using garbage values from uninitialized memory diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index cb549bfcd..9e66c2828 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -475,6 +475,7 @@ class ServerArgs: lora_backend: str = "csgmv" max_lora_chunk_size: Optional[int] = 16 experts_shared_outer_loras: Optional[bool] = None + lora_strict_loading: bool = False # Kernel backend attention_backend: Optional[str] = None @@ -4961,6 +4962,13 @@ class ServerArgs: "(expert_dim=1). Use --no-experts-shared-outer-loras to force disable. " "By default this is auto-detected from adapter weights.", ) + parser.add_argument( + "--lora-strict-loading", + default=ServerArgs.lora_strict_loading, + action=argparse.BooleanOptionalAction, + help="Enable strict loading for LoRA adapters. " + "When set, mismatched or missing keys in the adapter weights will raise an error.", + ) # Kernel backend parser.add_argument( @@ -6669,14 +6677,14 @@ class ServerArgs: "Expected a list or a dictionary." ) - # Expand target modules + # Normalize target modules to a set; keep {"all"} as a sentinel + # that gets resolved model-awarely in lora_manager.init_lora_shapes(). if self.lora_target_modules: self.lora_target_modules = set(self.lora_target_modules) if "all" in self.lora_target_modules: assert ( len(self.lora_target_modules) == 1 ), "If 'all' is specified in --lora-target-modules, it should be the only module specified." - self.lora_target_modules = set(SUPPORTED_LORA_TARGET_MODULES) # Ensure sufficient information is provided for LoRA initialization. assert self.lora_paths or (