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."""
|
||||
|
||||
Reference in New Issue
Block a user