Update LoRA Weights via Tensor (#16226)

Co-authored-by: PopSoda2002 <zhouhp.me@gmail.com>
This commit is contained in:
lg(x)
2026-01-10 17:36:43 +08:00
committed by GitHub
parent aeb480c11f
commit 3a8b44fe89
12 changed files with 628 additions and 31 deletions

View File

@@ -47,6 +47,7 @@ from sglang.srt.managers.io_struct import (
GenerateReqInput,
GetWeightsByNameReqInput,
InitWeightsUpdateGroupReqInput,
LoadLoRAAdapterFromTensorsReqInput,
LoadLoRAAdapterReqInput,
MultimodalDataInputFormat,
ReleaseMemoryOccupationReqInput,
@@ -600,6 +601,22 @@ class Engine(EngineBase):
self.tokenizer_manager.get_weights_by_name(obj, None)
)
def load_lora_adapter_from_tensors(
self, lora_name: str, tensors: List[Tuple[str, torch.Tensor]], config_dict: Dict
):
# Load LoRA adapter again
serialized_tensors = MultiprocessingSerializer.serialize(
tensors, output_str=True
)
lora_req = LoadLoRAAdapterFromTensorsReqInput(
lora_name=lora_name,
config_dict=config_dict,
serialized_tensors=serialized_tensors,
)
return self.loop.run_until_complete(
self.tokenizer_manager.load_lora_adapter_from_tensors(lora_req, None)
)
def load_lora_adapter(self, lora_name: str, lora_path: str, pinned: bool = False):
"""Load a new LoRA adapter without re-launching the engine."""

View File

@@ -103,6 +103,7 @@ from sglang.srt.managers.io_struct import (
GetWeightsByNameReqInput,
InitWeightsSendGroupForRemoteInstanceReqInput,
InitWeightsUpdateGroupReqInput,
LoadLoRAAdapterFromTensorsReqInput,
LoadLoRAAdapterReqInput,
OpenSessionReqInput,
ParseFunctionCallReq,
@@ -1062,6 +1063,21 @@ async def load_lora_adapter(obj: LoadLoRAAdapterReqInput, request: Request):
)
@app.api_route("/load_lora_adapter_from_tensors", methods=["POST"])
async def load_lora_adapter_from_tensors(
obj: LoadLoRAAdapterFromTensorsReqInput, request: Request
):
"""Load a new LoRA adapter from tensors without re-launching the server."""
result = await _global_state.tokenizer_manager.load_lora_adapter_from_tensors(
obj, request
)
if result.success:
return ORJSONResponse(result, status_code=HTTPStatus.OK)
else:
return ORJSONResponse(result, status_code=HTTPStatus.BAD_REQUEST)
@app.api_route("/unload_lora_adapter", methods=["POST"])
async def unload_lora_adapter(obj: UnloadLoRAAdapterReqInput, request: Request):
"""Load a new LoRA adapter without re-launching the server."""

View File

@@ -74,44 +74,55 @@ class LoRAAdapter(nn.Module):
self.embedding_layers: Dict[str, torch.Tensor] = {}
self.added_tokens_embeddings: Dict[str, torch.Tensor] = {}
# initialize the LoRA weights to cpu
def initialize_weights(self):
model_path = self.config.path
loader = DefaultModelLoader(self.load_config)
revision = getattr(self.config.hf_config, "revision", None)
# Get normalized target modules for filtering
for name, loaded_weight in loader._get_weights_iterator(
DefaultModelLoader.Source(
model_path, revision=revision, fall_back_to_pt=True
)
):
self._process_weight(name, loaded_weight)
self._normalize_weights()
def initialize_weights_from_tensors(self, tensors: Dict[str, torch.Tensor]):
for name, tensor in tensors.items():
self._process_weight(name, tensor)
self._normalize_weights()
def _process_weight(self, name: str, loaded_weight: torch.Tensor):
from sglang.srt.lora.utils import get_normalized_target_modules
normalized_target_modules = get_normalized_target_modules(
self.config.target_modules
)
for name, loaded_weight in loader._get_weights_iterator(
DefaultModelLoader.Source(
model_path, revision=revision, fall_back_to_pt=True
)
):
layer_id = get_layer_id(name)
if layer_id is not None:
self.layers[layer_id].weights[name] = loaded_weight.cpu()
elif "embed_tokens" in name or "lm_head" in name:
# Check if this module is declared in target_modules before loading
module_name = "embed_tokens" if "embed_tokens" in name else "lm_head"
if module_name in normalized_target_modules:
self.embedding_layers[name] = loaded_weight.cpu()
else:
logger.debug(
f"Skipping {name} as '{module_name}' is not in adapter's target_modules: {self.config.target_modules}"
)
elif "input_embeddings" in name or "output_embeddings" in name:
# added/extra token emb
self.added_tokens_embeddings[name] = loaded_weight.cpu()
assert loaded_weight.shape[0] == self.config.lora_added_tokens_size, (
f"LoRA adapter {self.uid} has extra_vocab_size {self.config.extra_vocab_size} specified in the config, "
f"but the loaded weight has {loaded_weight.shape[0]} extra vocab size"
layer_id = get_layer_id(name)
if layer_id is not None:
self.layers[layer_id].weights[name] = loaded_weight.cpu()
elif "embed_tokens" in name or "lm_head" in name:
# Check if this module is declared in target_modules before loading
module_name = "embed_tokens" if "embed_tokens" in name else "lm_head"
if module_name in normalized_target_modules:
self.embedding_layers[name] = loaded_weight.cpu()
else:
logger.debug(
f"Skipping {name} as '{module_name}' is not in adapter's target_modules: {self.config.target_modules}"
)
elif "input_embeddings" in name or "output_embeddings" in name:
# added/extra token emb
self.added_tokens_embeddings[name] = loaded_weight.cpu()
assert loaded_weight.shape[0] == self.config.lora_added_tokens_size, (
f"LoRA adapter {self.uid} has extra_vocab_size {self.config.extra_vocab_size} specified in the config, "
f"but the loaded weight has {loaded_weight.shape[0]} extra vocab size"
)
def _normalize_weights(self):
# normalize kv_proj and gate_up_proj
for layer in self.layers:
weight_names = list(layer.weights.keys())

View File

@@ -14,6 +14,7 @@
import json
import os
from typing import Dict, Optional
from huggingface_hub import snapshot_download
@@ -21,20 +22,34 @@ from huggingface_hub import snapshot_download
class LoRAConfig:
def __init__(
self,
path: str,
path: Optional[str] = None,
config_dict: Optional[Dict] = None,
added_tokens_config: Optional[Dict] = None,
) -> None:
self.path = path
self.hf_config = self.get_lora_config()
self.target_modules = self.hf_config["target_modules"]
if config_dict is not None:
self.hf_config = config_dict
self.added_tokens_config = added_tokens_config
else:
self.hf_config = self.get_lora_config()
self.added_tokens_config = self.get_added_tokens_config()
self.target_modules = self.hf_config["target_modules"]
self.r = self.hf_config["r"]
self.lora_alpha = self.hf_config["lora_alpha"]
self.added_tokens_config = self.get_added_tokens_config()
self.lora_added_tokens_size = (
len(self.added_tokens_config) if self.added_tokens_config is not None else 0
)
@classmethod
def from_dict(
cls,
config_dict: Dict,
added_tokens_config: Optional[Dict] = None,
) -> "LoRAConfig":
return cls(config_dict=config_dict, added_tokens_config=added_tokens_config)
def get_lora_config(self, dummy=False):
if dummy:
raise NotImplementedError()

View File

@@ -444,6 +444,56 @@ class LoRAManager:
lora_adapter.initialize_weights()
self.loras[lora_ref.lora_id] = lora_adapter
def load_lora_weights_from_tensors(
self, lora_ref: LoRARef, tensors: Dict[str, torch.Tensor]
):
"""
Load the weights of a LoRA adapter from tensors to CPU memory.
"""
lora_adapter = LoRAAdapter(
lora_ref.lora_id,
self.configs[lora_ref.lora_id],
self.base_hf_config,
self.load_config,
self.lora_backend,
)
lora_adapter.initialize_weights_from_tensors(tensors)
self.loras[lora_ref.lora_id] = lora_adapter
def load_lora_adapter_from_tensors(
self,
lora_ref: LoRARef,
tensors: Dict[str, torch.Tensor],
config_dict: Dict,
added_tokens_config: Optional[Dict] = None,
) -> LoRAUpdateOutput:
"""
Load a single LoRA adapter from tensors and config dict.
"""
assert (
lora_ref.lora_name is not None and lora_ref.lora_path is not None
), "LoRARef must have both lora_name and lora_path set for loading."
assert (
lora_ref.lora_id not in self.loras
), f"LoRA adapter with ID {lora_ref.lora_id} is already loaded. This should have been verified before request is sent to the backend."
try:
new_adapter = LoRAConfig.from_dict(config_dict, added_tokens_config)
self.validate_new_adapter(new_adapter, lora_ref)
self.configs[lora_ref.lora_id] = new_adapter
self.load_lora_weights_from_tensors(lora_ref, tensors)
self.lora_refs[lora_ref.lora_id] = lora_ref
self.num_pinned_loras += int(lora_ref.pinned)
except Exception as e:
return self.create_lora_update_result(
success=False,
error_message=str(e),
)
return self.create_lora_update_result(success=True)
def init_memory_pool(self):
"""(Re)initialize the LoRA memory pool based on the current configurations."""
self.memory_pool = LoRAMemoryPool(

View File

@@ -92,7 +92,9 @@ def get_hidden_dim(
# if contain extra tokens will be added; otherwise is 0.
return config.hidden_size, config.vocab_size + lora_added_vocab_size
else:
raise NotImplementedError()
raise NotImplementedError(
"get_hidden_dim not implemented for " + module_name
)
def get_normalized_target_modules(

View File

@@ -1643,6 +1643,24 @@ class UnloadLoRAAdapterReqInput(BaseReq):
)
@dataclass
class LoadLoRAAdapterFromTensorsReqInput(BaseReq):
lora_name: str
config_dict: Dict[str, Any]
serialized_tensors: str
pinned: bool = False
added_tokens_config: Optional[Dict[str, Any]] = None
lora_id: Optional[str] = None
def to_ref(self) -> LoRARef:
return LoRARef(
lora_id=self.lora_id,
lora_name=self.lora_name,
lora_path="__tensor__",
pinned=self.pinned,
)
@dataclass
class LoRAUpdateOutput(BaseReq):
success: bool
@@ -1650,7 +1668,9 @@ class LoRAUpdateOutput(BaseReq):
loaded_adapters: Optional[Dict[str, LoRARef]] = None
LoadLoRAAdapterReqOutput = UnloadLoRAAdapterReqOutput = LoRAUpdateOutput
LoadLoRAAdapterReqOutput = UnloadLoRAAdapterReqOutput = (
LoadLoRAAdapterFromTensorsReqOutput
) = LoRAUpdateOutput
class BlockReqType(Enum):

View File

@@ -92,6 +92,8 @@ from sglang.srt.managers.io_struct import (
InitWeightsSendGroupForRemoteInstanceReqInput,
InitWeightsSendGroupForRemoteInstanceReqOutput,
InitWeightsUpdateGroupReqInput,
LoadLoRAAdapterFromTensorsReqInput,
LoadLoRAAdapterFromTensorsReqOutput,
LoadLoRAAdapterReqInput,
LoadLoRAAdapterReqOutput,
OpenSessionReqInput,
@@ -1052,6 +1054,10 @@ class Scheduler(
(RpcReqInput, self.handle_rpc_request),
(ExpertDistributionReq, self.expert_distribution_handle),
(LoadLoRAAdapterReqInput, self.load_lora_adapter),
(
LoadLoRAAdapterFromTensorsReqInput,
self.load_lora_adapter_from_tensors,
),
(UnloadLoRAAdapterReqInput, self.unload_lora_adapter),
(GetLoadReqInput, self.get_load),
(PauseGenerationReqInput, self.pause_generation),
@@ -2703,6 +2709,14 @@ class Scheduler(
result = self.tp_worker.load_lora_adapter(recv_req)
return result
def load_lora_adapter_from_tensors(
self, recv_req: LoadLoRAAdapterFromTensorsReqInput
) -> LoadLoRAAdapterFromTensorsReqOutput:
"""In-place loading a new lora adapter from serialized tensors."""
result = self.tp_worker.load_lora_adapter_from_tensors(recv_req)
return result
def unload_lora_adapter(
self, recv_req: UnloadLoRAAdapterReqInput
) -> UnloadLoRAAdapterReqOutput:

View File

@@ -45,6 +45,8 @@ from sglang.srt.managers.io_struct import (
InitWeightsSendGroupForRemoteInstanceReqOutput,
InitWeightsUpdateGroupReqInput,
InitWeightsUpdateGroupReqOutput,
LoadLoRAAdapterFromTensorsReqInput,
LoadLoRAAdapterFromTensorsReqOutput,
LoadLoRAAdapterReqInput,
LoadLoRAAdapterReqOutput,
LoRAUpdateOutput,
@@ -617,6 +619,76 @@ class TokenizerCommunicatorMixin:
error_message=str(e),
)
async def load_lora_adapter_from_tensors(
self: TokenizerManager,
obj: LoadLoRAAdapterFromTensorsReqInput,
_: Optional[fastapi.Request] = None,
) -> LoadLoRAAdapterFromTensorsReqOutput:
self.auto_create_handle_loop()
try:
if not self.server_args.enable_lora:
raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
)
assert (
self.server_args.dp_size == 1
), "dp_size must be 1 for dynamic lora loading"
logger.info(
"Start load Lora adapter from tensors. Lora name=%s",
obj.lora_name,
)
async with self.lora_update_lock:
new_adapter = LoRARef(
lora_name=obj.lora_name,
lora_path="__tensor__",
pinned=obj.pinned,
)
obj.lora_id = new_adapter.lora_id
result = (await self.update_lora_adapter_communicator(obj))[0]
if result.success:
await self.lora_registry.register(new_adapter)
self.lora_ref_cache[obj.lora_name] = new_adapter
if self.server_args.max_loaded_loras is not None:
while (
self.lora_registry.num_registered_loras
> self.server_args.max_loaded_loras
):
lru_lora_name = await self.lora_registry.lru_lora_name(
exclude_pinned=True
)
if lru_lora_name is None:
raise ValueError(
"Didn't find any LoRA adapters when trying to evict LRU LoRA adapter. "
f"LoRA registry is: {self.lora_registry._registry}"
)
logger.info(
f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
f"max allowed: {self.server_args.max_loaded_loras})"
)
unload_result = await self._unload_lora_adapter_locked(
UnloadLoRAAdapterReqInput(lora_name=lru_lora_name)
)
if not unload_result.success:
raise ValueError(
f"Error while unloading LRU LoRA adapter '{lru_lora_name}': "
f"{unload_result.error_message}"
)
del result.loaded_adapters[lru_lora_name]
return result
except ValueError as e:
return LoadLoRAAdapterFromTensorsReqOutput(
success=False,
error_message=str(e),
)
async def unload_lora_adapter(
self: TokenizerManager,
obj: UnloadLoRAAdapterReqInput,

View File

@@ -26,6 +26,7 @@ from sglang.srt.managers.io_struct import (
GetWeightsByNameReqInput,
InitWeightsSendGroupForRemoteInstanceReqInput,
InitWeightsUpdateGroupReqInput,
LoadLoRAAdapterFromTensorsReqInput,
LoadLoRAAdapterReqInput,
SendWeightsToRemoteInstanceReqInput,
UnloadLoRAAdapterReqInput,
@@ -189,6 +190,20 @@ class BaseTpWorker(ABC):
result = self.model_runner.unload_lora_adapter(recv_req.to_ref())
return result
def load_lora_adapter_from_tensors(
self, recv_req: LoadLoRAAdapterFromTensorsReqInput
):
# The LoRA code handles TP sharding internally using slice_lora_a_weights
# and slice_lora_b_weights methods (see lora/layers.py:46-49, mem_pool.py:437-440).
tensors = MultiprocessingSerializer.deserialize(recv_req.serialized_tensors)
result = self.model_runner.load_lora_adapter_from_tensors(
recv_req.to_ref(),
tensors,
recv_req.config_dict,
recv_req.added_tokens_config,
)
return result
def can_run_lora_batch(self, lora_ids: list[str]) -> bool:
lora_ids_set = set(lora_ids) if isinstance(lora_ids, list) else lora_ids
return self.model_runner.lora_manager.validate_lora_batch(lora_ids_set)

View File

@@ -1434,6 +1434,16 @@ class ModelRunner(ModelRunnerKVCacheMixin):
return result
def load_lora_adapter_from_tensors(
self, lora_ref: LoRARef, tensors, config_dict, added_tokens_config=None
):
logger.info(f"LoRA adapter loading from tensors starts: {lora_ref}.")
result = self.lora_manager.load_lora_adapter_from_tensors(
lora_ref, tensors, config_dict, added_tokens_config
)
logger.info(f"LoRA adapter loading from tensors completes: {lora_ref}.")
return result
def unload_lora_adapter(self, lora_ref: LoRARef):
"""Unload a lora adapter that was previously loaded during initialization or dynamic loading."""