From e5750a572ccb6c9278987ef51e5253686441a53d Mon Sep 17 00:00:00 2001 From: Bruce Wu Date: Wed, 18 Mar 2026 13:48:03 -0700 Subject: [PATCH] Support TP for lora lm_head layer (#18511) Co-authored-by: Ethan (Yusheng) Su --- python/sglang/srt/lora/layers.py | 121 +++++++++++++++++++---------- python/sglang/srt/lora/mem_pool.py | 55 +++++++------ python/sglang/srt/lora/utils.py | 19 +++++ 3 files changed, 132 insertions(+), 63 deletions(-) diff --git a/python/sglang/srt/lora/layers.py b/python/sglang/srt/lora/layers.py index 5da492512..d483f2b35 100644 --- a/python/sglang/srt/lora/layers.py +++ b/python/sglang/srt/lora/layers.py @@ -21,7 +21,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) from sglang.srt.lora.backend.base_backend import BaseLoRABackend -from sglang.srt.lora.utils import LoRABatchInfo +from sglang.srt.lora.utils import LoRABatchInfo, get_lm_head_lora_b_shard_size class BaseLayerWithLoRA(nn.Module): @@ -68,6 +68,23 @@ class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA): self.embed_dim = base_layer.embedding_dim self.vocab_size = base_layer.org_vocab_size + # Embedding LoRA with TP > 1 keeps weights fully replicated + # (unsharded) on every rank. This works correctly because the + # base VocabParallelEmbedding all-reduces its output before the + # LoRA delta is added, but it means each rank holds the full + # LoRA A (rank, vocab_size) and LoRA B (embed_dim, rank) tensors, + # which may cause OOM on large vocabularies or high LoRA ranks. + # + # input_scattered mode (DeepSeek-v2 MLA) skips the base + # all-reduce, making the unsharded LoRA approach mathematically + # incorrect — a sharded LoRA kernel would be needed. + if hasattr(base_layer, "tp_size") and base_layer.tp_size > 1: + from sglang.srt.layers.communicator import get_attn_tp_context + + assert ( + not get_attn_tp_context().allow_input_scattered + ), "VocabParallelEmbeddingWithLoRA with TP > 1 under input_scattered mode (e.g., DeepSeek-v2 MLA with --enable-attn-tp-input-scattered) is not fully supported and may produce incorrect results. Consider disabling input_scattered or removing embed_tokens from LoRA target modules." + self.output_offset = torch.tensor( [0, self.embed_dim], dtype=torch.int32, @@ -186,33 +203,28 @@ class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA): return base_output def slice_lora_a_weights(self, A: torch.Tensor, tp_rank: int): - # For TP=1, no slicing needed - # LoRA A weights (rank, vocab_size) are not sliced for embedding - # For TP>1, Need to modify code in: sglang/python/sglang/srt/lora/mem_pool.py - # return A - if tp_rank > 1: - raise NotImplementedError( - f"VocabParallelEmbeddingWithLoRA does not support tensor parallelism > 1. " - f"Got tp_size={tp_rank}" - ) + # LoRA A weights (rank, vocab_size) are kept unsharded. + # Each rank does a full embedding lookup; the result is complete + # on every rank and added to the already all-reduced base output. + return A def slice_lora_b_weights(self, B: torch.Tensor, tp_rank: int): - # For TP=1, no slicing needed - # LoRA B weights (embedding_dim, rank) would be sliced along embedding dimension for TP>1 - # For TP>1, Need to modify code in: sglang/python/sglang/srt/lora/mem_pool.py - # return B - if tp_rank > 1: - raise NotImplementedError( - f"VocabParallelEmbeddingWithLoRA does not support tensor parallelism > 1. " - f"Got tp_size={tp_rank}" - ) + # LoRA B weights (embedding_dim, rank) are kept unsharded. + # The base embedding output is all-reduced (full embedding_dim), + # so LoRA B must also produce full embedding_dim. + return B class ParallelLMHeadWithLoRA(BaseLayerWithLoRA): """ - Parallel LM Head layer with LoRA support (simplified for TP=1). + Parallel LM Head layer with LoRA support. The LM head computes logits = hidden_states @ (W + B @ A)^T + + With TP > 1, lm_head is column-parallel: each rank holds + weight (vocab_size/tp_size, hidden_size) and produces a shard + of logits. LoRA A is kept unsharded (rank, hidden_size) while + LoRA B is sliced along the vocab dimension to (vocab_size/tp_size, rank). """ def __init__( @@ -224,11 +236,40 @@ class ParallelLMHeadWithLoRA(BaseLayerWithLoRA): self.weight = base_layer.weight self.embed_dim = base_layer.embedding_dim self.vocab_size = base_layer.org_vocab_size - self.output_offset = torch.tensor( - [0, self.vocab_size], - dtype=torch.int32, - device=next(base_layer.parameters()).device, - ) + + tp_size = base_layer.tp_size if hasattr(base_layer, "tp_size") else 1 + + # lm_head LoRA keeps A unsharded and shards B along the vocab + # dimension, matching the column-parallel base output. This is + # incompatible with input_scattered mode where the all-reduce is + # skipped. + if tp_size > 1: + from sglang.srt.layers.communicator import get_attn_tp_context + + if get_attn_tp_context().allow_input_scattered: + raise ValueError( + "ParallelLMHeadWithLoRA is not compatible with " + "input_scattered mode (e.g., DeepSeek-v2 MLA with " + "--enable-attn-tp-input-scattered). Please disable " + "input_scattered or remove lm_head from LoRA " + "target modules." + ) + + self.shard_vocab_size = get_lm_head_lora_b_shard_size( + self.vocab_size, + shard_indices=base_layer.shard_indices, + ) + self.output_offset = torch.tensor( + [0, self.shard_vocab_size], + dtype=torch.int32, + device=next(base_layer.parameters()).device, + ) + else: + self.output_offset = torch.tensor( + [0, self.vocab_size], + dtype=torch.int32, + device=next(base_layer.parameters()).device, + ) def set_lora_info( self, @@ -338,24 +379,22 @@ class ParallelLMHeadWithLoRA(BaseLayerWithLoRA): self.lora_backend._lm_head_pass_idx = None def slice_lora_a_weights(self, A: torch.Tensor, tp_rank: int): - # For TP=1, no slicing needed - # For TP>1, need to modify code in: sglang/python/sglang/srt/lora/mem_pool.py - # return A - if tp_rank > 1: - raise NotImplementedError( - f"ParallelLMHeadWithLoRA does not support tensor parallelism > 1. " - f"Got tp_size={tp_rank}" - ) + # LoRA A weights (rank, hidden_size) are kept unsharded. + # Each rank receives full hidden_states, so A operates on full input. + return A def slice_lora_b_weights(self, B: torch.Tensor, tp_rank: int): - # For TP=1, no slicing needed - # For TP>1, would slice along vocab dimension, need to modify code in: sglang/python/sglang/srt/lora/mem_pool.py - # return B - if tp_rank > 1: - raise NotImplementedError( - f"ParallelLMHeadWithLoRA does not support tensor parallelism > 1. " - f"Got tp_size={tp_rank}" - ) + # lm_head is column-parallel: each rank produces vocab_size/tp_size (shard_vocab_size) + # logits. LoRA B (vocab_size, rank) must be sliced along the vocab + # dimension to match the sharded base output. + # Uses the base layer's shard_indices for the actual vocab range on + # this rank, staying consistent with base model weight sharding. + tp_size = self.base_layer.tp_size if hasattr(self.base_layer, "tp_size") else 1 + if tp_size <= 1: + return B + start_idx = self.base_layer.shard_indices.org_vocab_start_index + end_idx = self.base_layer.shard_indices.org_vocab_end_index + return B[start_idx:end_idx, :] class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA): diff --git a/python/sglang/srt/lora/mem_pool.py b/python/sglang/srt/lora/mem_pool.py index ede8b51be..9f65b1b00 100644 --- a/python/sglang/srt/lora/mem_pool.py +++ b/python/sglang/srt/lora/mem_pool.py @@ -14,6 +14,7 @@ from sglang.srt.lora.utils import ( ROW_PARALLELISM_LINEAR_LORA_NAMES, LoRAType, get_hidden_dim, + get_lm_head_lora_b_shard_size, get_normalized_target_modules, get_stacked_multiply, get_target_module_name, @@ -99,6 +100,17 @@ class LoRAMemoryPool: EMPTY_SLOT ] * self.max_loras_per_batch + # Cache lm_head shard_indices from the base model so that buffer + # allocation uses the same sharding as the base ParallelLMHead layer. + self.lm_head_shard_indices = None + if "lm_head" in target_modules and tp_size > 1: + from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead + + for _, module in base_model.named_modules(): + if isinstance(module, ParallelLMHead): + self.lm_head_shard_indices = module.shard_indices + break + self.init_buffers(base_model) def can_support(self, config: Union[LoRAConfig, Iterable[LoRAConfig]]) -> bool: @@ -156,7 +168,8 @@ class LoRAMemoryPool: input_dim, _ = get_hidden_dim( module_name, self.base_hf_config, base_model, 0, self.lora_added_tokens_size ) - # Have not imp self.tp_size > 1 yet. + # Embedding LoRA A is kept unsharded (full vocab) across TP ranks. + # Each rank does a full lookup; no vocab-dimension splitting needed. return ( self.max_loras_per_batch, max_lora_dim, @@ -194,7 +207,13 @@ class LoRAMemoryPool: _, output_dim = get_hidden_dim( module_name, self.base_hf_config, base_model, 0, self.lora_added_tokens_size ) - # Have not imp self.tp_size > 1 yet. + # lm_head is column-parallel so B is sharded; embed_tokens B stays + # unsharded (base output is all-reduced to full embed_dim). + if module_name == "lm_head": + output_dim = get_lm_head_lora_b_shard_size( + output_dim, + shard_indices=self.lm_head_shard_indices, + ) return ( self.max_loras_per_batch, output_dim, @@ -298,8 +317,8 @@ class LoRAMemoryPool: lora_adapters: Dict[str, LoRAAdapter], lora_modules: List[Dict[str, BaseLayerWithLoRA]], lora_refs: Dict[str, LoRARef], - lora_embed_tokens_module: Dict[str, BaseLayerWithLoRA], - lora_lm_head_module: Dict[str, BaseLayerWithLoRA], + lora_embed_tokens_module: Optional[BaseLayerWithLoRA], + lora_lm_head_module: Optional[BaseLayerWithLoRA], ): def get_available_buffer_slot(): # 1. Prioritize empty slots @@ -379,8 +398,8 @@ class LoRAMemoryPool: buffer_id: int, lora_adapter: LoRAAdapter, lora_modules: List[Dict[str, BaseLayerWithLoRA]], - lora_embed_tokens_module: Dict[str, BaseLayerWithLoRA], - lora_lm_head_module: Dict[str, BaseLayerWithLoRA], + lora_embed_tokens_module: Optional[BaseLayerWithLoRA], + lora_lm_head_module: Optional[BaseLayerWithLoRA], ): def load_lora_weight_tensor( buffer_view: torch.Tensor, weight: Optional[torch.Tensor] @@ -487,13 +506,8 @@ class LoRAMemoryPool: and ("lora_embedding_B" in name or "lora_B" in name) ): lora_b_weights = weights - # [to-do] support TP - # if self.tp_size > 1: - # cur_module = lora_embeddings_modules[target_module] - # for module_name, module in cur_module: - # lora_b_weights = module.slice_lora_b_weights( - # lora_b_weights, self.tp_rank - # ) + # TP is supported by keeping embedding LoRA B unsharded; + # no slicing needed. buffer_view = self.embedding_B_buffer[target_module][ buffer_id, :, :lora_rank @@ -518,18 +532,15 @@ class LoRAMemoryPool: and ("lora_embedding_B" in name or "lora_B" in name) ): lora_b_weights = weights - # [to-do] support TP - # if self.tp_size > 1: - # cur_module = lora_embeddings_modules[target_module] - # for module_name, module in cur_module: - # lora_b_weights = module.slice_lora_b_weights( - # lora_b_weights, self.tp_rank - # ) + # Slice B along vocab dimension for this TP rank + if self.tp_size > 1 and lora_lm_head_module is not None: + lora_b_weights = lora_lm_head_module.slice_lora_b_weights( + lora_b_weights, self.tp_rank + ) buffer_view = self.lm_head_B_buffer[target_module][ - # buffer_id, :lora_rank, : org_vocab_size + extra_vocab_size buffer_id, - : (org_vocab_size + self.lora_added_tokens_size), + : lora_b_weights.shape[0], :lora_rank, ] load_lora_weight_tensor(buffer_view, lora_b_weights) diff --git a/python/sglang/srt/lora/utils.py b/python/sglang/srt/lora/utils.py index d121639ca..286fce5fa 100644 --- a/python/sglang/srt/lora/utils.py +++ b/python/sglang/srt/lora/utils.py @@ -171,6 +171,25 @@ EMBEDDING_NAMES = ["embed_tokens", "lm_head"] ROW_PARALLELISM_LINEAR_LORA_NAMES = ["o_proj", "down_proj"] +def get_lm_head_lora_b_shard_size(output_dim: int, shard_indices=None) -> int: + """Get the LoRA B output dimension for lm_head, accounting for TP. + + lm_head is column-parallel, so its LoRA B must be sharded along the + vocab dimension to match the base output. When shard_indices is + provided, the returned size reflects the base model's actual per-rank + vocab partition. + + Args: + output_dim: Full (unsharded) output dimension (vocab_size). + shard_indices: VocabParallelEmbeddingShardIndices from the base + ParallelLMHead layer. When provided, returns the per-rank + org vocab size from the base model's actual sharding. + """ + if shard_indices is not None: + return shard_indices.num_org_elements + return output_dim + + def generate_sequence_lengths( forward_batch: ForwardBatch, device: Optional[torch.device] = None ) -> torch.Tensor: