[Feature] add LoRADrainer to address high P99 TTFT (#17913)
This commit is contained in:
@@ -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 |
|
||||
|
||||
@@ -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>
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user