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
co-authored by PopSoda2002
parent aeb480c11f
commit 3a8b44fe89
12 changed files with 628 additions and 31 deletions
+17
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."""
@@ -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."""
+35 -24
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())
+20 -5
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()
+50
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(
+3 -1
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(
+21 -1
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):
+14
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:
@@ -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,
+15
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)
@@ -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."""
@@ -0,0 +1,355 @@
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=90, suite="stage-b-test-small-1-gpu")
register_amd_ci(est_time=90, suite="stage-b-test-small-1-gpu")
import json
import os
import unittest
import torch
from huggingface_hub import snapshot_download
from safetensors.torch import load_file
import sglang as sgl
from sglang.test.test_utils import CustomTestCase
MODEL_PATH = "Qwen/Qwen3-0.6B"
LORA_REPO = "charent/self_cognition_Alice"
TEST_PROMPT = "Hello, my name is"
EXPECTED_OUTPUT = (
" Alice, and I am a software engineer. I am excited to share my journey"
)
MAX_NEW_TOKENS = 16
class TestLoRALoadFromTensor(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.engine = sgl.Engine(
model_path=MODEL_PATH,
enable_lora=True,
max_lora_rank=64,
lora_target_modules=[
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
],
mem_fraction_static=0.6,
log_level="error",
)
lora_adapter = snapshot_download(
repo_id=LORA_REPO,
allow_patterns=["adapter_model.safetensors", "adapter_config.json"],
)
# Load tensors and config from downloaded adapter
cls.lora_tensors = load_file(
os.path.join(lora_adapter, "adapter_model.safetensors")
)
with open(os.path.join(lora_adapter, "adapter_config.json"), "r") as f:
cls.lora_config_dict = json.load(f)
def test_lora_lru_eviction(self):
print("[Test]Testing LRU LoRA eviction...")
MAX_LOADED_LORAS = 8
print(f"[Test]Max loaded LoRAs: {MAX_LOADED_LORAS}")
test_engine = sgl.Engine(
model_path=MODEL_PATH,
enable_lora=True,
max_lora_rank=64,
lora_target_modules=[
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
],
mem_fraction_static=0.6,
log_level="error",
max_loaded_loras=MAX_LOADED_LORAS,
)
# Load 10 LoRA adapters, max allowed is 8
# This should trigger LRU eviction when we exceed the limit
TEST_LORA_COUNT = 10
for i in range(TEST_LORA_COUNT):
print(f"[Test]Loading LoRA adapter {i+1}/10: self_cognition_Alice_{i}")
result = test_engine.load_lora_adapter_from_tensors(
lora_name=f"self_cognition_Alice_{i}",
tensors=self.lora_tensors,
config_dict=self.lora_config_dict,
)
self.assertTrue(
result.success,
f"Failed to load LoRA adapter {i}: {result.error_message}",
)
print(
f"[Test]Successfully loaded LoRA {i+1}, current loaded adapters: {list(result.loaded_adapters.keys())}"
)
EXPECTED_LORA_ADAPTERS = [
"self_cognition_Alice_2",
"self_cognition_Alice_3",
"self_cognition_Alice_4",
"self_cognition_Alice_5",
"self_cognition_Alice_6",
"self_cognition_Alice_7",
"self_cognition_Alice_8",
"self_cognition_Alice_9",
]
EXPECTED_LORA_COUNT = 8
self.assertEqual(
len(result.loaded_adapters),
EXPECTED_LORA_COUNT,
f"Loaded adapters count does not match expected result: {len(result.loaded_adapters)} != {EXPECTED_LORA_COUNT}",
)
self.assertEqual(
list(result.loaded_adapters.keys()),
EXPECTED_LORA_ADAPTERS,
f"Loaded adapters do not match expected result: {list(result.loaded_adapters.keys())} != {EXPECTED_LORA_ADAPTERS}",
)
print(
f"[Test]LRU eviction test passed! Final loaded adapters: {len(result.loaded_adapters)}"
)
def test_lora_e2e_load_from_tensor_params(self):
print("[Test]Testing LoRA load from tensor params...")
result = self.engine.load_lora_adapter_from_tensors(
lora_name="self_cognition_Alice",
tensors=self.lora_tensors,
config_dict=self.lora_config_dict,
)
self.assertTrue(
result.success,
f"Failed to load LoRA from tensors: {result.error_message}",
)
output_without_lora = self.engine.generate(
prompt=[TEST_PROMPT],
sampling_params={
"max_new_tokens": MAX_NEW_TOKENS,
"temperature": 0.0,
},
)
output_lora = self.engine.generate(
prompt=[TEST_PROMPT],
sampling_params={
"max_new_tokens": MAX_NEW_TOKENS,
"temperature": 0.0,
},
lora_path=["self_cognition_Alice"],
)
print(f"[Without LoRA] {output_without_lora[0]}")
print(f"[With LoRA] {output_lora[0]}")
self.assertNotEqual(
output_without_lora[0]["text"][: len(EXPECTED_OUTPUT)],
EXPECTED_OUTPUT,
"Output before applying LoRA should not match expected result",
)
self.assertEqual(
output_lora[0]["text"][: len(EXPECTED_OUTPUT)],
EXPECTED_OUTPUT,
"Output after applying LoRA does not match expected result",
)
def test_lora_load_unload_load_from_tensor_params(self):
print("[Test]Testing LoRA load, unload, load from tensor params...")
# Load LoRA adapter from tensors
result = self.engine.load_lora_adapter_from_tensors(
lora_name="self_cognition_Alice_multiple",
tensors=self.lora_tensors,
config_dict=self.lora_config_dict,
)
self.assertTrue(
result.success,
f"Failed to load LoRA from tensors: {result.error_message}",
)
# Unload LoRA adapter
result = self.engine.unload_lora_adapter("self_cognition_Alice_multiple")
self.assertTrue(
result.success, f"Failed to unload LoRA: {result.error_message}"
)
with self.assertRaises(ValueError) as context:
output_lora = self.engine.generate(
prompt=[TEST_PROMPT],
sampling_params={
"max_new_tokens": MAX_NEW_TOKENS,
"temperature": 0.0,
},
lora_path=["self_cognition_Alice_multiple"],
)
# Load LoRA adapter again
result_again = self.engine.load_lora_adapter_from_tensors(
lora_name="self_cognition_Alice_multiple",
tensors=self.lora_tensors,
config_dict=self.lora_config_dict,
)
self.assertTrue(
result_again.success,
f"Failed to load LoRA from tensors: {result_again.error_message}",
)
output_lora_loaded_again = self.engine.generate(
prompt=[TEST_PROMPT],
sampling_params={
"max_new_tokens": MAX_NEW_TOKENS,
"temperature": 0.0,
},
lora_path=["self_cognition_Alice_multiple"],
)
print(f"[With LoRA Loaded again] {output_lora_loaded_again[0]}")
self.assertEqual(
output_lora_loaded_again[0]["text"][: len(EXPECTED_OUTPUT)],
EXPECTED_OUTPUT,
"Output after applying LoRA does not match expected result",
)
def test_lora_logp_diff_with_huggingface(self):
"""
Test comparing SGLang and HuggingFace LoRA logprobs when loading LoRA from tensors.
This verifies that loading LoRA adapters from tensors produces consistent logprobs
with HuggingFace.
"""
from sglang.test.runners import HFRunner, SRTRunner
from sglang.test.test_utils import DEFAULT_PORT_FOR_SRT_TEST_RUNNER
print("[Test]Testing LoRA logprob difference with HuggingFace...")
lora_name = "self_cognition_Alice_logprob_test"
prompts = [TEST_PROMPT]
# Step 1: Run SGLang with LoRA loaded from tensors
print("[Test]Running SGLang with LoRA from tensors...")
with SRTRunner(
MODEL_PATH,
torch_dtype=torch.float16,
model_type="generation",
tp_size=1,
max_loras_per_batch=1,
lora_backend="triton",
disable_cuda_graph=False,
disable_radix_cache=True,
port=DEFAULT_PORT_FOR_SRT_TEST_RUNNER,
mem_fraction_static=0.6,
enable_lora=True,
max_lora_rank=64,
lora_target_modules=[
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
],
) as srt_runner:
result = srt_runner.engine.load_lora_adapter_from_tensors(
lora_name=lora_name,
tensors=self.lora_tensors,
config_dict=self.lora_config_dict,
)
self.assertTrue(
result.success,
f"Failed to load LoRA from tensors: {result.error_message}",
)
# Run inference with loaded LoRA
srt_outputs = srt_runner.forward(
prompts,
max_new_tokens=MAX_NEW_TOKENS,
lora_paths=[lora_name],
)
# Step 2: Run HuggingFace with LoRA
print("[Test]Running HuggingFace with LoRA...")
torch.cuda.empty_cache()
with HFRunner(
MODEL_PATH,
torch_dtype=torch.float16,
model_type="generation",
patch_model_do_sample_false=True,
) as hf_runner:
hf_outputs = hf_runner.forward(
prompts,
max_new_tokens=MAX_NEW_TOKENS,
lora_paths=[LORA_REPO],
)
# Step 3: Compare results
sglang_text = srt_outputs.output_strs[0]
hf_text = hf_outputs.output_strs[0]
print(f"[Text Output]")
print(f" SGLang: {sglang_text}")
print(f" HuggingFace: {hf_text}")
# Compare prefill (input) logprobs
sglang_prefill = torch.tensor(srt_outputs.top_input_logprobs[0])
hf_prefill = torch.tensor(hf_outputs.top_input_logprobs[0])
prefill_diff = torch.abs(sglang_prefill - hf_prefill)
prefill_max_diff = torch.max(prefill_diff).item()
prefill_mean_diff = torch.mean(prefill_diff).item()
print(f"\n[Prefill Logprob Comparison]")
print(f" Shape: {list(sglang_prefill.shape)}")
print(f" Max difference: {prefill_max_diff:.6e}")
print(f" Mean difference: {prefill_mean_diff:.6e}")
# Compare decode (output) logprobs
sglang_decode = torch.tensor(srt_outputs.top_output_logprobs[0])
hf_decode = torch.tensor(hf_outputs.top_output_logprobs[0])
decode_diff = torch.abs(sglang_decode - hf_decode)
decode_max_diff = torch.max(decode_diff).item()
decode_mean_diff = torch.mean(decode_diff).item()
print(f"\n[Decode Logprob Comparison]")
print(f" Shape: {list(sglang_decode.shape)}")
print(f" Max difference: {decode_max_diff:.6e}")
print(f" Mean difference: {decode_mean_diff:.6e}")
# Assert logprobs are close (threshold 1e-1)
LOGPROB_THRESHOLD = 1e-1
self.assertLess(
prefill_max_diff,
LOGPROB_THRESHOLD,
f"Prefill logprob max difference too large: {prefill_max_diff:.6e} > {LOGPROB_THRESHOLD:.0e}",
)
self.assertLess(
decode_max_diff,
LOGPROB_THRESHOLD,
f"Decode logprob max difference too large: {decode_max_diff:.6e} > {LOGPROB_THRESHOLD:.0e}",
)
# Verify text outputs match expected
self.assertEqual(
sglang_text[: len(EXPECTED_OUTPUT)],
EXPECTED_OUTPUT,
"SGLang output does not match expected result",
)
print("\n[Test]LoRA logprob comparison test passed!")
@classmethod
def tearDownClass(cls):
cls.engine.shutdown()
if __name__ == "__main__":
unittest.main()