[Feature] add LoRADrainer to address high P99 TTFT (#17913)
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DRAIN_SCHEDULE_TOLERANCE = 1.2
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdapterStats:
|
||||
num_waiting_reqs: int = 0
|
||||
max_wait_time_secs: float = 0.0
|
||||
max_remaining_tokens: int = 0
|
||||
is_draining_for: Optional[str] = None
|
||||
|
||||
def _reset_stats(self):
|
||||
self.num_waiting_reqs = 0
|
||||
self.max_wait_time_secs = 0.0
|
||||
self.max_remaining_tokens = 0
|
||||
|
||||
def is_starving(self, drain_wait_threshold: float):
|
||||
return (
|
||||
self.max_wait_time_secs > drain_wait_threshold and self.num_waiting_reqs > 0
|
||||
)
|
||||
|
||||
|
||||
class LoRADrainer:
|
||||
"""
|
||||
Drainer for LoRA requests that manages draining. It tracks:
|
||||
- Number of waiting requests per adapter
|
||||
- Maximum wait time for requests needing each adapter
|
||||
- Maximum number of tokens needed for running requests for each adapter
|
||||
"""
|
||||
|
||||
def __init__(self, max_loras_per_batch: int, max_wait_time_secs: float = 0.0):
|
||||
self.max_loras_per_batch = max_loras_per_batch
|
||||
self.max_wait_time_secs = max_wait_time_secs
|
||||
self.adapter_to_stats: Dict[Optional[str], AdapterStats] = defaultdict(
|
||||
AdapterStats
|
||||
)
|
||||
|
||||
def update_draining_state(
|
||||
self,
|
||||
waiting_queue: List[Req],
|
||||
running_reqs: List[Req],
|
||||
) -> None:
|
||||
"""
|
||||
Update LoRA drainer state based on current waiting queue and running requests.
|
||||
|
||||
This method updates adapter statistics, identifies starving adapters that need
|
||||
to be scheduled, and marks adapters for draining to make room for starving ones.
|
||||
"""
|
||||
self._update_adapter_stats(waiting_queue, running_reqs)
|
||||
self._update_draining_loras(running_reqs)
|
||||
self._update_fully_drained_loras(running_reqs)
|
||||
|
||||
def _update_adapter_stats(
|
||||
self,
|
||||
waiting_queue: List[Req],
|
||||
running_reqs: List[Req],
|
||||
) -> None:
|
||||
for stats in self.adapter_to_stats.values():
|
||||
stats._reset_stats()
|
||||
|
||||
for req in waiting_queue:
|
||||
stats = self.adapter_to_stats[req.lora_id]
|
||||
|
||||
stats.num_waiting_reqs += 1
|
||||
stats.max_wait_time_secs = max(
|
||||
stats.max_wait_time_secs,
|
||||
time.monotonic() - req.time_stats.wait_queue_entry_time,
|
||||
)
|
||||
|
||||
for req in running_reqs:
|
||||
stats = self.adapter_to_stats[req.lora_id]
|
||||
|
||||
stats.max_remaining_tokens = max(
|
||||
stats.max_remaining_tokens,
|
||||
req.sampling_params.max_new_tokens - len(req.output_ids),
|
||||
)
|
||||
|
||||
def _update_draining_loras(self, running_reqs: List[Req]) -> None:
|
||||
"""
|
||||
Select LoRA adapters to drain based on starvation detection.
|
||||
|
||||
This method identifies adapters in the waiting queue that are "starving"
|
||||
(waiting too long) and marks currently running adapters as "draining"
|
||||
to make room for the starving adapters. Draining adapters will not
|
||||
accept new requests, allowing them to complete and free up slots.
|
||||
"""
|
||||
running_adapter_ids = {req.lora_id for req in running_reqs}
|
||||
if len(running_adapter_ids) < self.max_loras_per_batch:
|
||||
return None
|
||||
|
||||
starving_adapters = set()
|
||||
draining_for_adapters = set()
|
||||
for adapter_id, stats in self.adapter_to_stats.items():
|
||||
if stats.is_starving(self.max_wait_time_secs):
|
||||
starving_adapters.add(adapter_id)
|
||||
|
||||
draining_for_adapter = stats.is_draining_for
|
||||
if draining_for_adapter is not None:
|
||||
draining_for_adapters.add(draining_for_adapter)
|
||||
|
||||
new_starving_adapters = starving_adapters - draining_for_adapters
|
||||
if not new_starving_adapters:
|
||||
return None
|
||||
|
||||
sorted_new_starving_adapters = sorted(
|
||||
new_starving_adapters,
|
||||
key=lambda adapter: self.adapter_to_stats[adapter].max_wait_time_secs,
|
||||
reverse=True,
|
||||
)
|
||||
|
||||
eligible_to_drain_adapters = {
|
||||
adapter
|
||||
for adapter in running_adapter_ids
|
||||
if self.adapter_to_stats[adapter].is_draining_for is None
|
||||
}
|
||||
|
||||
for starving_adapter in sorted_new_starving_adapters:
|
||||
if not eligible_to_drain_adapters:
|
||||
break
|
||||
|
||||
min_eligible_adapter = min(
|
||||
eligible_to_drain_adapters,
|
||||
key=lambda adapter_id: self.adapter_to_stats[
|
||||
adapter_id
|
||||
].max_remaining_tokens,
|
||||
)
|
||||
|
||||
self.adapter_to_stats[min_eligible_adapter].is_draining_for = (
|
||||
starving_adapter
|
||||
)
|
||||
logger.debug(
|
||||
f"LoRA adapter {min_eligible_adapter} is draining for {starving_adapter}"
|
||||
)
|
||||
|
||||
eligible_to_drain_adapters.remove(min_eligible_adapter)
|
||||
|
||||
def _update_fully_drained_loras(self, running_reqs: List[Req]) -> None:
|
||||
"""
|
||||
Clear draining state for adapters that have fully drained.
|
||||
|
||||
An adapter is considered fully drained when it was marked as draining
|
||||
but no longer has any running requests.
|
||||
"""
|
||||
running_adapter_ids = {req.lora_id for req in running_reqs}
|
||||
for adapter_id, stats in self.adapter_to_stats.items():
|
||||
if stats.is_draining_for is None:
|
||||
continue
|
||||
|
||||
if adapter_id not in running_adapter_ids:
|
||||
logger.debug(f"LoRA adapter {adapter_id} finished draining")
|
||||
stats.is_draining_for = None
|
||||
|
||||
def can_schedule(self, req: Req) -> bool:
|
||||
"""
|
||||
Check if a request can be scheduled based on draining state.
|
||||
|
||||
If the adapter for this request is currently draining, only allow
|
||||
scheduling if the request's max_new_tokens is within tolerance of
|
||||
the max remaining tokens for the draining adapter.
|
||||
"""
|
||||
stats = self.adapter_to_stats[req.lora_id]
|
||||
if not stats.is_draining_for:
|
||||
return True
|
||||
|
||||
return (
|
||||
req.sampling_params.max_new_tokens
|
||||
<= stats.max_remaining_tokens * DRAIN_SCHEDULE_TOLERANCE
|
||||
)
|
||||
@@ -78,6 +78,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_drainer import LoRADrainer
|
||||
from sglang.srt.lora.lora_overlap_loader import LoRAOverlapLoader
|
||||
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
|
||||
from sglang.srt.managers.io_struct import (
|
||||
@@ -473,6 +474,15 @@ class Scheduler(
|
||||
# Init request dispatcher
|
||||
self.init_request_dispatcher()
|
||||
|
||||
# Init LoRA drainer for fair scheduling
|
||||
if self.server_args.lora_drain_wait_threshold > 0.0:
|
||||
self.lora_drainer = LoRADrainer(
|
||||
server_args.max_loras_per_batch,
|
||||
server_args.lora_drain_wait_threshold,
|
||||
)
|
||||
else:
|
||||
self.lora_drainer = None
|
||||
|
||||
# Init LoRA overlap loader
|
||||
if self.enable_lora_overlap_loading:
|
||||
self.lora_overlap_loader = LoRAOverlapLoader(
|
||||
@@ -2639,23 +2649,16 @@ class Scheduler(
|
||||
if self.enable_lora:
|
||||
running_loras = {req.lora_id for req in self.running_batch.reqs}
|
||||
|
||||
if self.lora_drainer:
|
||||
self.lora_drainer.update_draining_state(
|
||||
self.waiting_queue,
|
||||
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 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
|
||||
if self.enable_lora and not self._can_schedule_lora_req(req, running_loras):
|
||||
continue
|
||||
|
||||
running_bs = len(self.running_batch.reqs)
|
||||
if len(adder.can_run_list) >= self.get_num_allocatable_reqs(running_bs):
|
||||
@@ -2800,6 +2803,34 @@ class Scheduler(
|
||||
|
||||
return new_batch
|
||||
|
||||
def _can_schedule_lora_req(
|
||||
self, req: Req, running_loras: set[Optional[str]]
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a LoRA request can be scheduled.
|
||||
|
||||
This method checks two conditions:
|
||||
1. The drainer allows scheduling (based on draining state)
|
||||
2. The LoRA adapter can be loaded (either already running or can be added)
|
||||
"""
|
||||
if self.lora_drainer and not self.lora_drainer.can_schedule(req):
|
||||
return False
|
||||
|
||||
if req.lora_id in running_loras:
|
||||
return True
|
||||
|
||||
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
|
||||
return self.lora_overlap_loader.try_overlap_load_lora(
|
||||
req.lora_id, running_loras
|
||||
)
|
||||
else:
|
||||
new_lora_set = {req.lora_id} | running_loras
|
||||
return self.tp_worker.model_runner.lora_manager.validate_lora_batch(
|
||||
new_lora_set
|
||||
)
|
||||
|
||||
def update_running_batch(self, batch: ScheduleBatch) -> Optional[ScheduleBatch]:
|
||||
"""Update the current running decoding batch."""
|
||||
initial_bs = batch.batch_size()
|
||||
|
||||
@@ -492,6 +492,7 @@ class ServerArgs:
|
||||
experts_shared_outer_loras: Optional[bool] = None
|
||||
lora_use_virtual_experts: bool = False
|
||||
lora_strict_loading: bool = False
|
||||
lora_drain_wait_threshold: float = 0.0
|
||||
|
||||
# Kernel backend
|
||||
attention_backend: Optional[str] = None
|
||||
@@ -5258,6 +5259,12 @@ class ServerArgs:
|
||||
help="Enable strict loading for LoRA adapters. "
|
||||
"When set, mismatched or missing keys in the adapter weights will raise an error.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora-drain-wait-threshold",
|
||||
type=float,
|
||||
default=ServerArgs.lora_drain_wait_threshold,
|
||||
help="When any LoRA adapter request waits longer than this threshold (in seconds), the scheduler will selectively drain one running adapter to make room. This mitigates extreme tail latency under high or skewed workloads by preventing a small set of adapters from monopolizing batch slots. Set to 0 to disable draining (default).",
|
||||
)
|
||||
|
||||
# Kernel backend
|
||||
parser.add_argument(
|
||||
@@ -7086,6 +7093,10 @@ class ServerArgs:
|
||||
)
|
||||
logger.info("Virtual expert computation enabled.")
|
||||
|
||||
assert (
|
||||
self.lora_drain_wait_threshold >= 0.0
|
||||
), "--lora-drain-wait-threshold must be non-negative."
|
||||
|
||||
def validate_buckets_rule(self, arg_name: str, buckets_rule: List[str]):
|
||||
if not buckets_rule:
|
||||
return
|
||||
|
||||
@@ -822,6 +822,7 @@ def run_lora_batch_splitting_equivalence_test(
|
||||
disable_cuda_graph: bool = True,
|
||||
disable_radix_cache: bool = True,
|
||||
enable_lora_overlap_loading: Optional[bool] = None,
|
||||
lora_drain_wait_threshold: float = 0.0,
|
||||
):
|
||||
"""
|
||||
Test that SRT correctly handles batch splitting with multiple LoRA adapters.
|
||||
@@ -839,6 +840,9 @@ def run_lora_batch_splitting_equivalence_test(
|
||||
attention_backend: Attention backend to use
|
||||
disable_cuda_graph: Whether to disable CUDA graph
|
||||
disable_radix_cache: Whether to disable radix cache
|
||||
lora_drain_wait_threshold: When any LoRA adapter request waits longer than
|
||||
this threshold (in seconds), the scheduler will selectively drain one
|
||||
running adapter to make room. Set to 0 to disable draining (default).
|
||||
"""
|
||||
max_loras_per_batch = 2
|
||||
|
||||
@@ -851,9 +855,14 @@ def run_lora_batch_splitting_equivalence_test(
|
||||
max_new_tokens = 64
|
||||
base_path = model_case.base
|
||||
|
||||
maybe_drain_info = (
|
||||
f", lora_drain_wait_threshold={lora_drain_wait_threshold}"
|
||||
if lora_drain_wait_threshold > 0
|
||||
else ""
|
||||
)
|
||||
print(
|
||||
f"\n========== Testing batch splitting on base '{base_path}', "
|
||||
f"dtype={torch_dtype} =========="
|
||||
f"dtype={torch_dtype}{maybe_drain_info} =========="
|
||||
)
|
||||
|
||||
prompts = [TEST_MULTIPLE_BATCH_PROMPTS[0]] * 3
|
||||
@@ -897,6 +906,7 @@ def run_lora_batch_splitting_equivalence_test(
|
||||
attention_backend=attention_backend,
|
||||
disable_cuda_graph=disable_cuda_graph,
|
||||
disable_radix_cache=disable_radix_cache,
|
||||
lora_drain_wait_threshold=lora_drain_wait_threshold,
|
||||
) as srt_runner:
|
||||
for batch_idx, (batch_prompts, lora_paths) in enumerate(test_cases):
|
||||
print(f"\n--- Batch {batch_idx + 1} ---")
|
||||
|
||||
@@ -587,6 +587,7 @@ class SRTRunner:
|
||||
json_model_override_args: Optional[dict[str, Any]] = None,
|
||||
lora_eviction_policy: str = "lru",
|
||||
enable_deterministic_inference: bool = False,
|
||||
lora_drain_wait_threshold: float = 0.0,
|
||||
):
|
||||
self.model_type = model_type
|
||||
self.is_generation = model_type == "generation"
|
||||
@@ -648,6 +649,7 @@ class SRTRunner:
|
||||
),
|
||||
lora_eviction_policy=lora_eviction_policy,
|
||||
enable_deterministic_inference=enable_deterministic_inference,
|
||||
lora_drain_wait_threshold=lora_drain_wait_threshold,
|
||||
**spec_kwargs,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user