Support TP for lora lm_head layer (#18511)

Co-authored-by: Ethan (Yusheng) Su <yushengsu.thu@gmail.com>
This commit is contained in:
Bruce Wu
2026-03-18 13:48:03 -07:00
committed by GitHub
parent 8f0f36c64b
commit e5750a572c
3 changed files with 132 additions and 63 deletions

View File

@@ -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):

View File

@@ -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)

View File

@@ -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: