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."""