add documentation example for LoRA overlap loading and cleanup unused function (#17464)

This commit is contained in:
Glen Liu
2026-01-24 02:33:16 -05:00
committed by GitHub
parent 3992a023e6
commit a6280b2a23
5 changed files with 140 additions and 16 deletions

View File

@@ -1,4 +1,5 @@
import logging
from typing import Type
from sglang.srt.lora.backend.base_backend import BaseLoRABackend
@@ -50,7 +51,7 @@ def create_flashinfer_backend():
)
def get_backend_from_name(name: str) -> BaseLoRABackend:
def get_backend_from_name(name: str) -> Type[BaseLoRABackend]:
"""
Get corresponding backend class from backend's name
"""

View File

@@ -55,13 +55,13 @@ class LoRAManager:
max_loras_per_batch: int,
load_config: LoadConfig,
dtype: torch.dtype,
server_args: ServerArgs,
lora_backend: str = "triton",
tp_size: int = 1,
tp_rank: int = 0,
max_lora_rank: Optional[int] = None,
target_modules: Optional[Iterable[str]] = None,
lora_paths: Optional[List[LoRARef]] = None,
server_args: Optional[ServerArgs] = None,
):
self.base_model: torch.nn.Module = base_model
self.base_hf_config: AutoConfig = base_hf_config

View File

@@ -195,10 +195,6 @@ class BaseTpWorker(ABC):
)
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)
def forward_batch_embedding(self, model_worker_batch: ModelWorkerBatch):
forward_batch = ForwardBatch.init_new(model_worker_batch, self.model_runner)
logits_output = self.model_runner.forward(forward_batch).logits_output

View File

@@ -1424,13 +1424,13 @@ class ModelRunner(ModelRunnerKVCacheMixin):
max_loras_per_batch=self.server_args.max_loras_per_batch,
load_config=self.load_config,
dtype=self.dtype,
server_args=self.server_args,
lora_backend=self.server_args.lora_backend,
tp_size=self.tp_size,
tp_rank=self.tp_rank,
max_lora_rank=self.server_args.max_lora_rank,
target_modules=self.server_args.lora_target_modules,
lora_paths=self.server_args.lora_paths,
server_args=self.server_args,
)
def load_lora_adapter(self, lora_ref: LoRARef):