From 76b9c8de6f495bf738e61de01ff8be0e74e99cbc Mon Sep 17 00:00:00 2001
From: Glen Liu <62917497+glenliu21@users.noreply.github.com>
Date: Sat, 2 May 2026 19:13:43 -0400
Subject: [PATCH] [Feature] add LoRADrainer to address high P99 TTFT (#17913)
---
docs/advanced_features/lora.ipynb | 2 +
docs/advanced_features/server_arguments.md | 1 +
docs_new/docs/advanced_features/lora.mdx | 4 +-
.../advanced_features/server_arguments.mdx | 6 +
python/sglang/srt/lora/lora_drainer.py | 191 ++++++++++++++++++
python/sglang/srt/managers/scheduler.py | 61 ++++--
python/sglang/srt/server_args.py | 11 +
python/sglang/test/lora_utils.py | 12 +-
python/sglang/test/runners.py | 2 +
test/registered/lora/test_lora_drainer.py | 131 ++++++++++++
10 files changed, 404 insertions(+), 17 deletions(-)
create mode 100644 python/sglang/srt/lora/lora_drainer.py
create mode 100644 test/registered/lora/test_lora_drainer.py
diff --git a/docs/advanced_features/lora.ipynb b/docs/advanced_features/lora.ipynb
index 8e6e6d0a0..230bd700f 100644
--- a/docs/advanced_features/lora.ipynb
+++ b/docs/advanced_features/lora.ipynb
@@ -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."
diff --git a/docs/advanced_features/server_arguments.md b/docs/advanced_features/server_arguments.md
index 7675b4bbb..8ad1c0881 100644
--- a/docs/advanced_features/server_arguments.md
+++ b/docs/advanced_features/server_arguments.md
@@ -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 |
diff --git a/docs_new/docs/advanced_features/lora.mdx b/docs_new/docs/advanced_features/lora.mdx
index 7e0f5d95b..3ed6b4430 100644
--- a/docs_new/docs/advanced_features/lora.mdx
+++ b/docs_new/docs/advanced_features/lora.mdx
@@ -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.
diff --git a/docs_new/docs/advanced_features/server_arguments.mdx b/docs_new/docs/advanced_features/server_arguments.mdx
index 4285ff521..902601797 100644
--- a/docs_new/docs/advanced_features/server_arguments.mdx
+++ b/docs_new/docs/advanced_features/server_arguments.mdx
@@ -1126,6 +1126,12 @@ Please consult the documentation below and [server_args.py](https://github.com/s
`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` |
+ Type: float |
+
diff --git a/python/sglang/srt/lora/lora_drainer.py b/python/sglang/srt/lora/lora_drainer.py
new file mode 100644
index 000000000..d60c5787a
--- /dev/null
+++ b/python/sglang/srt/lora/lora_drainer.py
@@ -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
+ )
diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py
index 2c3d98647..c1f79bb0d 100644
--- a/python/sglang/srt/managers/scheduler.py
+++ b/python/sglang/srt/managers/scheduler.py
@@ -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()
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index c5bba8351..ec4c4a279 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -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
diff --git a/python/sglang/test/lora_utils.py b/python/sglang/test/lora_utils.py
index 9ce95e7b8..39258ea18 100644
--- a/python/sglang/test/lora_utils.py
+++ b/python/sglang/test/lora_utils.py
@@ -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} ---")
diff --git a/python/sglang/test/runners.py b/python/sglang/test/runners.py
index c2d84ff2f..3e8f3f6fd 100644
--- a/python/sglang/test/runners.py
+++ b/python/sglang/test/runners.py
@@ -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,
)
diff --git a/test/registered/lora/test_lora_drainer.py b/test/registered/lora/test_lora_drainer.py
new file mode 100644
index 000000000..5b7d97e6d
--- /dev/null
+++ b/test/registered/lora/test_lora_drainer.py
@@ -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()