Update LoRA Weights via Tensor (#16226)
Co-authored-by: PopSoda2002 <zhouhp.me@gmail.com>
This commit is contained in:
@@ -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."""
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user