[Feature] add LoRADrainer to address high P99 TTFT (#17913)

This commit is contained in:
Glen Liu
2026-05-02 16:13:43 -07:00
committed by GitHub
parent 200944b415
commit 76b9c8de6f
10 changed files with 404 additions and 17 deletions
+2
View File
@@ -47,6 +47,8 @@
"\n",
"* `--max-lora-chunk-size`: Maximum chunk size for the ChunkedSGMV LoRA backend. Only used when --lora-backend is 'csgmv'. Choosing a larger value might improve performance. Please tune this value based on your hardware and workload as needed. Defaults to 16.\n",
"\n",
"* `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. 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).\n",
"\n",
"* `tp_size`: LoRA serving along with Tensor Parallelism is supported by SGLang. `tp_size` controls the number of GPUs for tensor parallelism. More details on the tensor sharding strategy can be found in [S-Lora](https://arxiv.org/pdf/2311.03285) paper.\n",
"\n",
"From client side, the user needs to provide a list of strings as input batch, and a list of adaptor names that each input sequence corresponds to."
@@ -258,6 +258,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s
| `--lora-eviction-policy` | LoRA adapter eviction policy when the GPU memory pool is full. | `lru` | `lru`, `fifo` |
| `--lora-backend` | Choose the kernel backend for multi-LoRA serving. | `csgmv` | `triton`, `csgmv`, `ascend`, `torch_native` |
| `--max-lora-chunk-size` | Maximum chunk size for the ChunkedSGMV LoRA backend. Only used when `--lora-backend` is `csgmv`. Larger values may improve performance. | `16` | `16`, `32`, `64`, `128` |
| `--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. 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). | `0.0` | Type: float |
## Kernel Backends (Attention, Sampling, Grammar, GEMM)
| Argument | Description | Defaults | Options |
+3 -1
View File
@@ -29,7 +29,9 @@ The following server arguments are relevant for multi-LoRA serving:
* `lora_target_modules`: The union set of all target modules where LoRA should be applied (e.g., `q_proj`, `k_proj`, `gate_proj`). If not specified, it will be automatically inferred from the adapters provided in `--lora-paths`. This argument is needed when you expect to dynamically load adapters of different target modules after server startup. You can also set it to `all` to enable LoRA for all supported modules. However, enabling LoRA on additional modules introduces a minor performance overhead. If your application is performance-sensitive, we recommend only specifying the modules for which you plan to load adapters.
* `--max-lora-chunk-size`: Maximum chunk size for the ChunkedSGMV LoRA backend. Only used when --lora-backend is 'csgmv'. Choosing a larger value might improve performance. Please tune this value based on your hardware and workload as needed. Defaults to 16.
* `max_lora_chunk_size`: Maximum chunk size for the ChunkedSGMV LoRA backend. Only used when --lora-backend is 'csgmv'. Choosing a larger value might improve performance. Please tune this value based on your hardware and workload as needed. Defaults to 16.
* `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. 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).
* `tp_size`: LoRA serving along with Tensor Parallelism is supported by SGLang. `tp_size` controls the number of GPUs for tensor parallelism. More details on the tensor sharding strategy can be found in [S-Lora](https://arxiv.org/pdf/2311.03285) paper.
@@ -1126,6 +1126,12 @@ Please consult the documentation below and [server_args.py](https://github.com/s
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`16`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>16</code>, <code>32</code>, <code>64</code>, <code>128</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--lora-drain-wait-threshold`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>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).</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`0`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Type: float</td>
</tr>
</tbody>
</table>
+191
View File
@@ -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
)
+46 -15
View File
@@ -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()
+11
View File
@@ -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
+11 -1
View File
@@ -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} ---")
+2
View File
@@ -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,
)
+131
View File
@@ -0,0 +1,131 @@
import unittest
from types import SimpleNamespace
from typing import cast
from unittest import mock
from sglang.srt.lora.lora_drainer import LoRADrainer
from sglang.srt.managers.schedule_batch import Req
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.lora_utils import (
CI_MULTI_LORA_MODELS,
run_lora_batch_splitting_equivalence_test,
)
from sglang.test.test_utils import is_in_ci
register_cuda_ci(est_time=100, suite="stage-b-test-1-gpu-small")
register_amd_ci(est_time=100, suite="stage-b-test-1-gpu-small-amd")
MOCK_START_TIME = 1000.0
LORA_DRAIN_WAIT_THRESHOLD = 3.0
def make_req(lora_id, wait_queue_entry_time, max_new_tokens, output_len=0):
time_stats = SimpleNamespace(wait_queue_entry_time=wait_queue_entry_time)
sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens)
req_ns = SimpleNamespace(
lora_id=lora_id,
time_stats=time_stats,
sampling_params=sampling_params,
output_ids=[0] * output_len,
)
return cast(Req, req_ns)
class TestLoRADrainer(unittest.TestCase):
def test_update_draining_marks_adapter(self):
if is_in_ci():
return
with mock.patch("time.monotonic", return_value=MOCK_START_TIME):
drainer = LoRADrainer(
max_loras_per_batch=1, max_wait_time_secs=LORA_DRAIN_WAIT_THRESHOLD
)
# Waiting request for adapter 'A' that has been waiting longer than threshold
wait_entry = MOCK_START_TIME - (LORA_DRAIN_WAIT_THRESHOLD + 0.01)
waiting_req = make_req("A", wait_entry, max_new_tokens=10)
running_req = make_req("B", wait_entry, max_new_tokens=100, output_len=0)
drainer.update_draining_state(
waiting_queue=[waiting_req],
running_reqs=[running_req],
)
# Running adapter 'B' should be marked as draining for 'A'
self.assertEqual(drainer.adapter_to_stats["B"].is_draining_for, "A")
# Once running adapter 'B' finishes running, it should no longer be draining
drainer.update_draining_state(waiting_queue=[waiting_req], running_reqs=[])
self.assertIsNone(drainer.adapter_to_stats["B"].is_draining_for)
with mock.patch("time.monotonic", return_value=MOCK_START_TIME):
drainer = LoRADrainer(
max_loras_per_batch=2, max_wait_time_secs=LORA_DRAIN_WAIT_THRESHOLD
)
# Two starving adapters should cause two running adapters to drain.
wait_entryA = MOCK_START_TIME - (LORA_DRAIN_WAIT_THRESHOLD + 0.05)
wait_entryD = MOCK_START_TIME - (LORA_DRAIN_WAIT_THRESHOLD + 0.01)
starving_a = make_req("A", wait_entryA, max_new_tokens=10)
starving_d = make_req("D", wait_entryD, max_new_tokens=10)
# Running adapters B and C with different remaining tokens
running_b = make_req("B", wait_entryA, max_new_tokens=5, output_len=0)
running_c = make_req("C", wait_entryA, max_new_tokens=100, output_len=0)
drainer.update_draining_state(
waiting_queue=[starving_a, starving_d],
running_reqs=[running_b, running_c],
)
# B (smaller remaining tokens) should be drained for the most-starved adapter 'A'
self.assertEqual(drainer.adapter_to_stats["B"].is_draining_for, "A")
# C should be drained for the other starving adapter 'D'
self.assertEqual(drainer.adapter_to_stats["C"].is_draining_for, "D")
def test_can_schedule_respects_draining_tolerance(self):
if is_in_ci():
return
with mock.patch("time.monotonic", return_value=MOCK_START_TIME):
drainer = LoRADrainer(
max_loras_per_batch=1, max_wait_time_secs=LORA_DRAIN_WAIT_THRESHOLD
)
wait_entry = MOCK_START_TIME - (LORA_DRAIN_WAIT_THRESHOLD + 0.01)
starving_req = make_req("A", wait_entry, max_new_tokens=10)
running_b = make_req("B", wait_entry, max_new_tokens=15, output_len=0)
drainer.update_draining_state(
waiting_queue=[starving_req],
running_reqs=[running_b],
)
self.assertEqual(drainer.adapter_to_stats["B"].is_draining_for, "A")
# max_new_tokens is less than running adapter B's max_new_tokens
req_ok = make_req(
lora_id="B", wait_queue_entry_time=0, max_new_tokens=10, output_len=0
)
self.assertTrue(drainer.can_schedule(req_ok))
# max_new_tokens is more than running adapter B's max_new_tokens
req_bad = make_req(
lora_id="B", wait_queue_entry_time=0, max_new_tokens=20, output_len=0
)
self.assertFalse(drainer.can_schedule(req_bad))
def test_batch_splitting_with_drainer(self):
run_lora_batch_splitting_equivalence_test(
model_cases=CI_MULTI_LORA_MODELS,
attention_backend="torch_native",
disable_cuda_graph=True,
disable_radix_cache=True,
lora_drain_wait_threshold=LORA_DRAIN_WAIT_THRESHOLD,
)
if __name__ == "__main__":
unittest.main()