[Feature] Add LoRA support for embedding layers (#14177)
Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com> Co-authored-by: Beichen-Ma <bm685@cornell.edu>
This commit is contained in:
co-authored by
Baizhou Zhang
Beichen-Ma
parent
9ad02b799d
commit
0c63fb9420
@@ -1,4 +1,7 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
@@ -13,8 +16,12 @@ from sglang.srt.layers.linear import (
|
||||
QKVParallelLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
ParallelLMHead,
|
||||
VocabParallelEmbedding,
|
||||
)
|
||||
from sglang.srt.lora.backend.base_backend import BaseLoRABackend
|
||||
from sglang.srt.lora.utils import LoRABatchInfo
|
||||
|
||||
|
||||
class BaseLayerWithLoRA(nn.Module):
|
||||
@@ -45,11 +52,10 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
|
||||
class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA):
|
||||
"""
|
||||
Vocab parallel embedding layer with support for LoRA (Low-Rank Adaptation).
|
||||
Vocab parallel embedding layer with LoRA support (simplified for TP=1, no extra tokens).
|
||||
|
||||
Note: The current version does not yet implement the LoRA functionality.
|
||||
This class behaves exactly the same as the base VocabParallelEmbedding.
|
||||
Future versions will integrate LoRA functionality to support efficient parameter fine-tuning.
|
||||
For embedding layers: output = base_embedding(x) + lora_B @ lora_A[x]
|
||||
where lora_A[x] is direct embedding lookup from lora_A weights.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -59,6 +65,237 @@ class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA):
|
||||
) -> None:
|
||||
super().__init__(base_layer, lora_backend)
|
||||
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.embed_dim],
|
||||
dtype=torch.int32,
|
||||
device=next(base_layer.parameters()).device,
|
||||
)
|
||||
|
||||
def set_lora_info(
|
||||
self,
|
||||
new_embeddings_buffer: Optional[torch.Tensor], # For extra tokens
|
||||
embedding_A_buffer: torch.Tensor,
|
||||
embedding_B_buffer: torch.Tensor,
|
||||
):
|
||||
"""Set LoRA buffers for embedding layer."""
|
||||
self.set_lora = True
|
||||
self.new_embeddings_buffer = new_embeddings_buffer
|
||||
self.embedding_A_buffer = embedding_A_buffer # (num_loras, rank, vocab_size)
|
||||
self.embedding_B_buffer = embedding_B_buffer # (num_loras, embed_dim, rank)
|
||||
|
||||
def apply_lora(
|
||||
self, base_output: torch.Tensor, input_: torch.Tensor, batch_info
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply LoRA to base embedding output.
|
||||
Formula: output = base_output + lora_B @ lora_A_embedding(input_)
|
||||
"""
|
||||
|
||||
# Efficient embedding lookup for LoRA A (already support extra token embedding process)
|
||||
lora_a_output = self.run_lora_a_embedding(input_, batch_info)
|
||||
|
||||
# Apply LoRA B weights using backend
|
||||
lora_output = self.lora_backend.run_lora_b_sgemm(
|
||||
x=lora_a_output,
|
||||
weights=self.embedding_B_buffer,
|
||||
output_offset=self.output_offset,
|
||||
base_output=base_output,
|
||||
)
|
||||
return lora_output
|
||||
|
||||
def run_lora_a_embedding(
|
||||
self, input_: torch.Tensor, batch_info: LoRABatchInfo
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply LoRA A weights using efficient embedding lookup with CUDA graph support.
|
||||
Maps tokens to their corresponding LoRA adapters internally.
|
||||
It also includes added/extra token processing.
|
||||
"""
|
||||
# Efficient embedding lookup for LoRA A (already support extra token embedding process)
|
||||
lora_a_output = self.lora_backend.run_lora_a_embedding(
|
||||
input_ids=input_,
|
||||
weights=self.embedding_A_buffer,
|
||||
vocab_size=self.vocab_size,
|
||||
extra_embeddings=(
|
||||
self.new_embeddings_buffer
|
||||
if hasattr(self, "new_embeddings_buffer")
|
||||
and self.new_embeddings_buffer is not None
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
return lora_a_output
|
||||
|
||||
def extra_token_embedding(
|
||||
self, input_: torch.Tensor, base_output: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Need to impl:
|
||||
|
||||
Process extra tokens (tokens >= vocab_size) by looking up their embeddings
|
||||
from the new_embeddings_buffer and replacing them in base_output.
|
||||
|
||||
Args:
|
||||
input_: (s,) token IDs
|
||||
base_output: (s, embed_dim) base embedding output to be modified in-place
|
||||
|
||||
Returns:
|
||||
base_output: (s, embed_dim) modified input base_output (tensor[0,0,0,...]) with extra token embeddings
|
||||
"""
|
||||
# return base_output
|
||||
raise NotImplementedError(
|
||||
"Error in sglang/python/sglang/srt/lora/layers.py - VocabParallelEmbeddingWithLoRA \n"
|
||||
"Current SGLang codebase did not support tuned lora with extra/added tokens. \n"
|
||||
"[TODO]: \n"
|
||||
"1. Refer to this commit: https://github.com/yushengsu-thu/sglang/commit/90415211eee8a28a316de262583d4d33fa615d10#diff-191177438bcc223837963de63c005850371f8c8a860acb153b26744b66ecc623 to complete \n"
|
||||
"2. And then you need to modified the en/decoder tokenizer - tokenizer_manager.py to support extra_token_embedding in-place. \n"
|
||||
)
|
||||
|
||||
def forward(self, input_: torch.Tensor):
|
||||
"""
|
||||
Forward pass with LoRA support and CUDA graph compatibility.
|
||||
|
||||
Extra tokens (tokens >= vocab_size) are now handled efficiently
|
||||
in the backend's run_lora_a_embedding method.
|
||||
"""
|
||||
batch_info = self.lora_backend.batch_info
|
||||
|
||||
# Get base embedding output
|
||||
# For tokens >= vocab_size, base_layer will clamp or handle them
|
||||
# We mask them to 0 to avoid out-of-bounds access
|
||||
added_tokens_mask = input_ > self.vocab_size - 1
|
||||
base_output = self.base_layer.forward(input_.masked_fill(added_tokens_mask, 0))
|
||||
|
||||
# [TODO] SGLang did not support extra/added token process; thus, self.extra_token_embedding only return original input_ now
|
||||
# Extra tokens - It will replace extra token embedding with self.new_embeddings_buffer's emb (Default is 0)
|
||||
if (
|
||||
hasattr(self, "new_embeddings_buffer")
|
||||
and self.new_embeddings_buffer is not None
|
||||
):
|
||||
base_output = self.extra_token_embedding(input_, base_output)
|
||||
|
||||
# Apply LoRA if configured
|
||||
if self.set_lora:
|
||||
# The backend's run_lora_a_embedding now handles both regular
|
||||
# and extra tokens efficiently with CUDA graph support
|
||||
base_output = self.apply_lora(base_output, input_, batch_info)
|
||||
|
||||
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}"
|
||||
)
|
||||
|
||||
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}"
|
||||
)
|
||||
|
||||
|
||||
class ParallelLMHeadWithLoRA(BaseLayerWithLoRA):
|
||||
"""
|
||||
Parallel LM Head layer with LoRA support (simplified for TP=1).
|
||||
|
||||
The LM head computes logits = hidden_states @ (W + B @ A)^T
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_layer: ParallelLMHead,
|
||||
lora_backend: BaseLoRABackend,
|
||||
) -> None:
|
||||
super().__init__(base_layer, lora_backend)
|
||||
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,
|
||||
)
|
||||
|
||||
def set_lora_info(
|
||||
self,
|
||||
lm_head_A_buffer: torch.Tensor,
|
||||
lm_head_B_buffer: torch.Tensor,
|
||||
):
|
||||
"""Set LoRA buffers for LM head layer."""
|
||||
self.set_lora = True
|
||||
self.lm_head_A_buffer = lm_head_A_buffer # (num_loras, rank, hidden_dim)
|
||||
self.lm_head_B_buffer = lm_head_B_buffer # (num_loras, vocab_size, rank)
|
||||
|
||||
def apply_lora(
|
||||
self, base_output: torch.Tensor, hidden_states: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply LoRA to LM head layer.
|
||||
|
||||
For LM head: output = hidden @ (W + B @ A)^T
|
||||
= hidden @ W^T + hidden @ A^T @ B^T
|
||||
= base_output + (hidden @ A^T) @ B^T
|
||||
"""
|
||||
# Apply lora_A^T: hidden_states @ A^T
|
||||
lora_a_output = self.lora_backend.run_lora_a_sgemm(
|
||||
hidden_states, self.lm_head_A_buffer
|
||||
)
|
||||
|
||||
# Apply lora_B^T: lora_a_output @ B^T
|
||||
lora_output = self.lora_backend.run_lora_b_sgemm(
|
||||
x=lora_a_output,
|
||||
weights=self.lm_head_B_buffer,
|
||||
output_offset=self.output_offset,
|
||||
base_output=base_output,
|
||||
)
|
||||
|
||||
return lora_output
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor):
|
||||
# Apply base linear transformation
|
||||
base_output = F.linear(
|
||||
hidden_states, self.weight, bias=getattr(self.base_layer, "bias", None)
|
||||
)
|
||||
|
||||
# Apply LoRA if set
|
||||
if self.set_lora:
|
||||
base_output = self.apply_lora(base_output, hidden_states)
|
||||
|
||||
return base_output
|
||||
|
||||
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}"
|
||||
)
|
||||
|
||||
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}"
|
||||
)
|
||||
|
||||
|
||||
class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
|
||||
@@ -224,6 +461,7 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
|
||||
output_offset_cpu=self.output_offset_cpu,
|
||||
max_qkv_out_dim=self.max_qkv_out_dim,
|
||||
)
|
||||
|
||||
return lora_output
|
||||
|
||||
def slice_lora_a_weights(self, A: torch.Tensor, tp_rank: int):
|
||||
@@ -343,6 +581,7 @@ def get_lora_layer(
|
||||
) -> BaseLayerWithLoRA:
|
||||
supported_layer_types = {
|
||||
# the order matters
|
||||
ParallelLMHead: ParallelLMHeadWithLoRA,
|
||||
VocabParallelEmbedding: VocabParallelEmbeddingWithLoRA,
|
||||
QKVParallelLinear: QKVParallelLinearWithLoRA,
|
||||
MergedColumnParallelLinear: MergedColumnParallelLinearWithLoRA,
|
||||
|
||||
Reference in New Issue
Block a user