[Feature] overlap LoRA weight loading with compute (#15512)
This commit is contained in:
@@ -208,3 +208,14 @@ class LoRAAdapter(nn.Module):
|
||||
if "lora_A" in weight_name:
|
||||
weights[gate_up_name] = weights[gate_up_name].repeat(2, 1)
|
||||
# else: no-op as LoRA B weight is already stacked.
|
||||
|
||||
def pin_weights_in_cpu(self):
|
||||
for layer in self.layers:
|
||||
for name, weight in layer.weights.items():
|
||||
layer.weights[name] = weight.pin_memory()
|
||||
|
||||
for name, weight in self.embedding_layers.items():
|
||||
self.embedding_layers[name] = weight.pin_memory()
|
||||
|
||||
for name, weight in self.added_tokens_embeddings.items():
|
||||
self.added_tokens_embeddings[name] = weight.pin_memory()
|
||||
|
||||
@@ -72,6 +72,9 @@ class LoRAManager:
|
||||
self.tp_size: int = tp_size
|
||||
self.tp_rank: int = tp_rank
|
||||
self.lora_added_tokens_size: Optional[int] = None
|
||||
self.enable_lora_overlap_loading: Optional[bool] = (
|
||||
server_args.enable_lora_overlap_loading
|
||||
)
|
||||
|
||||
# Store eviction policy from server args
|
||||
self.eviction_policy = server_args.lora_eviction_policy
|
||||
@@ -208,7 +211,7 @@ class LoRAManager:
|
||||
|
||||
return self.create_lora_update_result(success=True)
|
||||
|
||||
def validate_lora_batch(self, lora_ids: set[str]) -> bool:
|
||||
def validate_lora_batch(self, lora_ids: set[Optional[str]]) -> bool:
|
||||
"""
|
||||
Validate if the LoRA IDs in the batch can be loaded into the current LoRA memory pool.
|
||||
"""
|
||||
@@ -239,9 +242,11 @@ class LoRAManager:
|
||||
|
||||
return required_slots <= mem_pool_vacancy
|
||||
|
||||
def prepare_lora_batch(self, forward_batch: ForwardBatch):
|
||||
def fetch_new_loras(
|
||||
self, new_loras: set[Optional[str]], running_loras: set[Optional[str]] = set()
|
||||
):
|
||||
# Load active loras into lora memory pool
|
||||
cur_uids = set(forward_batch.lora_ids)
|
||||
cur_uids = new_loras | running_loras
|
||||
|
||||
assert len(cur_uids) <= self.max_loras_per_batch
|
||||
self.memory_pool.prepare_lora_batch(
|
||||
@@ -253,6 +258,7 @@ class LoRAManager:
|
||||
lora_lm_head_module=self.lm_head_module, # merge into embedding or lora module
|
||||
)
|
||||
|
||||
def prepare_lora_batch(self, forward_batch: ForwardBatch):
|
||||
# set up batch info shared by all lora modules
|
||||
bs = forward_batch.batch_size
|
||||
|
||||
@@ -442,6 +448,11 @@ class LoRAManager:
|
||||
self.lora_backend,
|
||||
)
|
||||
lora_adapter.initialize_weights()
|
||||
|
||||
# If we want to overlap loading LoRA adapters with compute, they must be pinned in CPU memory
|
||||
if self.enable_lora_overlap_loading:
|
||||
lora_adapter.pin_weights_in_cpu()
|
||||
|
||||
self.loras[lora_ref.lora_id] = lora_adapter
|
||||
|
||||
def load_lora_weights_from_tensors(
|
||||
@@ -509,6 +520,9 @@ class LoRAManager:
|
||||
lora_added_tokens_size=self.lora_added_tokens_size,
|
||||
)
|
||||
|
||||
# Initializing memory pool with base model
|
||||
self.fetch_new_loras({None})
|
||||
|
||||
def set_lora_module(self, module_name, module):
|
||||
lora_module = get_lora_layer(module, self.lora_backend)
|
||||
replace_submodule(self.base_model, module_name, lora_module)
|
||||
|
||||
82
python/sglang/srt/lora/lora_overlap_loader.py
Normal file
82
python/sglang/srt/lora/lora_overlap_loader.py
Normal file
@@ -0,0 +1,82 @@
|
||||
import logging
|
||||
from enum import Enum, auto
|
||||
from typing import Dict, Optional
|
||||
|
||||
import torch
|
||||
from torch.cuda import Event as CudaEvent
|
||||
from torch.cuda import Stream as CudaStream
|
||||
from torch.cuda import StreamContext as CudaStreamContext
|
||||
|
||||
from sglang.srt.lora.lora_manager import LoRAManager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LoRAOverlapLoadStatus(Enum):
|
||||
LOADED = auto()
|
||||
LOADING = auto()
|
||||
NOT_LOADED = auto()
|
||||
|
||||
|
||||
class LoRAOverlapLoader:
|
||||
def __init__(self, lora_manager):
|
||||
self.lora_manager: LoRAManager = lora_manager
|
||||
self.device_module = torch.get_device_module(self.lora_manager.device)
|
||||
self.load_stream: CudaStream = self.device_module.Stream()
|
||||
self.load_stream_context: CudaStreamContext = self.device_module.stream(
|
||||
self.load_stream
|
||||
)
|
||||
self.lora_to_overlap_load_event: Dict[Optional[str], CudaEvent] = {}
|
||||
|
||||
def try_overlap_load_lora(
|
||||
self, lora_id: Optional[str], running_loras: set[Optional[str]]
|
||||
) -> bool:
|
||||
"""
|
||||
Check a LoRA adapter's asynchronous load status, and try to load it if there's capacity
|
||||
in the memory pool. Returns whether or not the adapter has been loaded.
|
||||
"""
|
||||
lora_pipeline_load_status = self._check_overlap_load_status(lora_id)
|
||||
if lora_pipeline_load_status == LoRAOverlapLoadStatus.LOADING:
|
||||
return False
|
||||
elif lora_pipeline_load_status == LoRAOverlapLoadStatus.NOT_LOADED:
|
||||
res = self._try_start_overlap_load(lora_id, running_loras)
|
||||
if res:
|
||||
logger.debug(f"Loading LoRA adapter {lora_id} asynchronously")
|
||||
|
||||
return False
|
||||
else:
|
||||
assert lora_pipeline_load_status == LoRAOverlapLoadStatus.LOADED
|
||||
return True
|
||||
|
||||
def _check_overlap_load_status(
|
||||
self, lora_id: Optional[str]
|
||||
) -> LoRAOverlapLoadStatus:
|
||||
if lora_id not in self.lora_to_overlap_load_event:
|
||||
return LoRAOverlapLoadStatus.NOT_LOADED
|
||||
|
||||
event = self.lora_to_overlap_load_event[lora_id]
|
||||
|
||||
if not event.query():
|
||||
return LoRAOverlapLoadStatus.LOADING
|
||||
|
||||
torch.cuda.current_stream().wait_event(event)
|
||||
del self.lora_to_overlap_load_event[lora_id]
|
||||
|
||||
return LoRAOverlapLoadStatus.LOADED
|
||||
|
||||
def _try_start_overlap_load(
|
||||
self, lora_id: Optional[str], running_loras: set[Optional[str]]
|
||||
) -> bool:
|
||||
loras_to_be_loaded = running_loras | self.lora_to_overlap_load_event.keys()
|
||||
|
||||
new_lora_set = {lora_id} | loras_to_be_loaded
|
||||
if not self.lora_manager.validate_lora_batch(new_lora_set):
|
||||
return False
|
||||
|
||||
with self.load_stream_context:
|
||||
self.lora_manager.fetch_new_loras({lora_id}, loras_to_be_loaded)
|
||||
event = self.device_module.Event()
|
||||
event.record(self.load_stream)
|
||||
|
||||
self.lora_to_overlap_load_event[lora_id] = event
|
||||
return True
|
||||
@@ -391,7 +391,7 @@ class LoRAMemoryPool:
|
||||
assert (
|
||||
buffer_view.shape == weight.shape
|
||||
), f"LoRA buffer shape {buffer_view.shape} does not match weight shape {weight.shape}."
|
||||
buffer_view.copy_(weight)
|
||||
buffer_view.copy_(weight, non_blocking=True)
|
||||
|
||||
if uid is None:
|
||||
for i in range(self.num_layer):
|
||||
|
||||
@@ -67,6 +67,7 @@ from sglang.srt.layers.dp_attention import (
|
||||
from sglang.srt.layers.moe import initialize_moe_config
|
||||
from sglang.srt.layers.quantization.fp4_utils import initialize_fp4_gemm_config
|
||||
from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
|
||||
from sglang.srt.lora.lora_overlap_loader import LoRAOverlapLoader
|
||||
from sglang.srt.managers.io_struct import (
|
||||
AbortReq,
|
||||
BaseBatchReq,
|
||||
@@ -287,6 +288,7 @@ class Scheduler(
|
||||
server_args.priority_scheduling_preemption_threshold
|
||||
)
|
||||
self.enable_lora = server_args.enable_lora
|
||||
self.enable_lora_overlap_loading = server_args.enable_lora_overlap_loading
|
||||
self.max_loras_per_batch = server_args.max_loras_per_batch
|
||||
self.enable_overlap = not server_args.disable_overlap_schedule
|
||||
self.enable_pdmux = server_args.enable_pdmux
|
||||
@@ -377,6 +379,12 @@ class Scheduler(
|
||||
# Init request dispatcher
|
||||
self.init_request_dispatcher()
|
||||
|
||||
# Init LoRA overlap loader
|
||||
if self.enable_lora_overlap_loading:
|
||||
self.lora_overlap_loader = LoRAOverlapLoader(
|
||||
self.tp_worker.model_runner.lora_manager
|
||||
)
|
||||
|
||||
# Init the grammar backend for constrained generation
|
||||
self.grammar_manager = GrammarManager(self)
|
||||
|
||||
@@ -1968,23 +1976,25 @@ class Scheduler(
|
||||
self.chunked_req = adder.add_chunked_req(self.chunked_req)
|
||||
|
||||
if self.enable_lora:
|
||||
lora_set = set([req.lora_id for req in self.running_batch.reqs])
|
||||
running_loras = {req.lora_id for req in self.running_batch.reqs}
|
||||
|
||||
# Get requests from the waiting queue to a new prefill batch
|
||||
for req in self.waiting_queue:
|
||||
|
||||
if self.enable_lora:
|
||||
new_lora_set = (
|
||||
lora_set
|
||||
| set([req.lora_id for req in adder.can_run_list])
|
||||
| set([req.lora_id])
|
||||
)
|
||||
if not self.tp_worker.can_run_lora_batch(new_lora_set):
|
||||
# Batch would exceed the LoRA slot limit.
|
||||
# Skip this request and try scheduling it in a future iteration.
|
||||
# Note: When eviction is needed, the eviction policy prefers to
|
||||
# evict LoRA adapters over base model (None) - see mem_pool.py.
|
||||
continue
|
||||
if self.enable_lora and req.lora_id not in running_loras:
|
||||
if self.enable_lora_overlap_loading:
|
||||
# For overlapping loading of LoRA weights with computation, we will load each adapter one at a time,
|
||||
# as opposed to loading them in one batch
|
||||
res = self.lora_overlap_loader.try_overlap_load_lora(
|
||||
req.lora_id, running_loras
|
||||
)
|
||||
if not res:
|
||||
continue
|
||||
else:
|
||||
new_lora_set = {req.lora_id} | running_loras
|
||||
if not self.tp_worker.model_runner.lora_manager.validate_lora_batch(
|
||||
new_lora_set
|
||||
):
|
||||
continue
|
||||
|
||||
running_bs = len(self.running_batch.reqs)
|
||||
if len(adder.can_run_list) >= self.get_num_allocatable_reqs(running_bs):
|
||||
@@ -2016,6 +2026,9 @@ class Scheduler(
|
||||
truncation_align_size=self.truncation_align_size,
|
||||
)
|
||||
|
||||
if self.enable_lora:
|
||||
running_loras.add(req.lora_id)
|
||||
|
||||
if res != AddReqResult.CONTINUE:
|
||||
if res == AddReqResult.NO_TOKEN:
|
||||
if self.enable_hierarchical_cache:
|
||||
|
||||
@@ -518,6 +518,11 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
|
||||
# Init lora information
|
||||
if model_runner.server_args.enable_lora:
|
||||
# In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool
|
||||
# as a batch, right before running the batch
|
||||
if not model_runner.server_args.enable_lora_overlap_loading:
|
||||
model_runner.lora_manager.fetch_new_loras(set(ret.lora_ids))
|
||||
|
||||
model_runner.lora_manager.prepare_lora_batch(ret)
|
||||
|
||||
return ret
|
||||
|
||||
@@ -415,6 +415,7 @@ class ServerArgs:
|
||||
|
||||
# LoRA
|
||||
enable_lora: Optional[bool] = None
|
||||
enable_lora_overlap_loading: Optional[bool] = None
|
||||
max_lora_rank: Optional[int] = None
|
||||
lora_target_modules: Optional[Union[set[str], List[str]]] = None
|
||||
lora_paths: Optional[
|
||||
@@ -3421,6 +3422,12 @@ class ServerArgs:
|
||||
action="store_true",
|
||||
help="Enable LoRA support for the model. This argument is automatically set to True if `--lora-paths` is provided for backward compatibility.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable-lora-overlap-loading",
|
||||
default=ServerArgs.enable_lora_overlap_loading,
|
||||
action="store_true",
|
||||
help="Enable asynchronous LoRA weight loading in order to overlap H2D transfers with GPU compute. This should be enabled if you find that your LoRA workloads are bottlenecked by adapter weight loading, for example when frequently loading large LoRA adapters.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-lora-rank",
|
||||
default=ServerArgs.max_lora_rank,
|
||||
@@ -4969,6 +4976,20 @@ class ServerArgs:
|
||||
)
|
||||
|
||||
if self.enable_lora:
|
||||
if self.enable_lora_overlap_loading is None:
|
||||
self.enable_lora_overlap_loading = False
|
||||
|
||||
if self.enable_lora_overlap_loading:
|
||||
# TODO (glenliu21): use some sort of buffer with eviction instead of enforcing a limit
|
||||
max_loaded_loras_limit = self.max_loras_per_batch * 2
|
||||
assert (
|
||||
self.max_loaded_loras is not None
|
||||
and self.max_loaded_loras <= max_loaded_loras_limit
|
||||
), (
|
||||
"Enabling LoRA overlap loading requires pinning LoRA adapter weights in CPU memory, "
|
||||
f"so --max-loaded-loras must be less than or equal to double --max-loras-per-batch: {max_loaded_loras_limit}"
|
||||
)
|
||||
|
||||
# Validate compatibility with speculative decoding
|
||||
if self.speculative_algorithm not in ["NGRAM", None]:
|
||||
raise ValueError(
|
||||
|
||||
@@ -96,6 +96,7 @@ CI_MULTI_LORA_MODELS = [
|
||||
),
|
||||
],
|
||||
max_loras_per_batch=2,
|
||||
max_loaded_loras=4,
|
||||
),
|
||||
]
|
||||
|
||||
@@ -285,6 +286,7 @@ def run_lora_test_one_by_one(
|
||||
torch_dtype: torch.dtype,
|
||||
max_new_tokens: int,
|
||||
backend: str = "csgmv",
|
||||
enable_lora_overlap_loading: Optional[bool] = None,
|
||||
disable_cuda_graph: bool = False,
|
||||
disable_radix_cache: bool = False,
|
||||
mem_fraction_static: float = 0.88,
|
||||
@@ -331,6 +333,7 @@ def run_lora_test_one_by_one(
|
||||
lora_paths=[
|
||||
adaptor.name for adaptor in model_case.adaptors if adaptor.name is not None
|
||||
],
|
||||
enable_lora_overlap_loading=enable_lora_overlap_loading,
|
||||
max_loras_per_batch=model_case.max_loras_per_batch,
|
||||
max_loaded_loras=model_case.max_loaded_loras,
|
||||
lora_backend=backend,
|
||||
|
||||
@@ -552,6 +552,7 @@ class SRTRunner:
|
||||
max_lora_rank: Optional[int] = None,
|
||||
lora_target_modules: Optional[List[str]] = None,
|
||||
enable_lora: Optional[bool] = None,
|
||||
enable_lora_overlap_loading: Optional[bool] = None,
|
||||
max_loaded_loras: Optional[int] = None,
|
||||
json_model_override_args: Optional[dict[str, Any]] = None,
|
||||
lora_eviction_policy: str = "lru",
|
||||
@@ -612,6 +613,7 @@ class SRTRunner:
|
||||
max_lora_rank=max_lora_rank,
|
||||
lora_target_modules=lora_target_modules,
|
||||
enable_lora=enable_lora,
|
||||
enable_lora_overlap_loading=enable_lora_overlap_loading,
|
||||
max_loaded_loras=max_loaded_loras,
|
||||
json_model_override_args=(
|
||||
json.dumps(json_model_override_args)
|
||||
|
||||
Reference in New Issue
Block a user