[DLLM] Basic dLLM scheduling strategy and implementation (#17484)
Signed-off-by: Zehuan Li <lizehuan.lzh@antgroup.com>
This commit is contained in:
@@ -0,0 +1,67 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import enum
|
||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
from sglang.srt.dllm.config import DllmConfig
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
|
|
||||||
|
|
||||||
|
class DllmReqPhase(str, enum.Enum):
|
||||||
|
STAGING_PREFILL = "staging_prefill"
|
||||||
|
STAGING_DECODE = "staging_decode"
|
||||||
|
INCOMING_PREFILL = "incoming_prefill"
|
||||||
|
INCOMING_DECODE = "incoming_decode"
|
||||||
|
|
||||||
|
|
||||||
|
class ReqDllmMixin:
|
||||||
|
def init_diffusion_llm(self: Req, dllm_config: DllmConfig):
|
||||||
|
self.dllm_phase: Optional[DllmReqPhase] = None
|
||||||
|
self.dllm_ids = []
|
||||||
|
self.dllm_block_offset = 0
|
||||||
|
self.dllm_config = dllm_config
|
||||||
|
|
||||||
|
if self.dllm_config is not None:
|
||||||
|
if len(self.origin_input_ids) < self.dllm_config.block_size:
|
||||||
|
self.dllm_phase = DllmReqPhase.INCOMING_DECODE
|
||||||
|
else:
|
||||||
|
self.dllm_phase = DllmReqPhase.INCOMING_PREFILL
|
||||||
|
|
||||||
|
def is_dllm(self: Req) -> bool:
|
||||||
|
return self.dllm_config is not None
|
||||||
|
|
||||||
|
def is_dllm_prefill(self: Req) -> bool:
|
||||||
|
return self.dllm_phase in [
|
||||||
|
DllmReqPhase.STAGING_PREFILL,
|
||||||
|
DllmReqPhase.INCOMING_PREFILL,
|
||||||
|
]
|
||||||
|
|
||||||
|
def determine_dllm_phase(self: Req):
|
||||||
|
prefix_length = len(self.prefix_indices)
|
||||||
|
min_required_length = prefix_length + self.dllm_config.block_size
|
||||||
|
|
||||||
|
if len(self.fill_ids) < min_required_length:
|
||||||
|
# still incoming stage
|
||||||
|
return
|
||||||
|
|
||||||
|
input_block = self.fill_ids[prefix_length:min_required_length]
|
||||||
|
is_prefill_phase = self.dllm_config.mask_id not in input_block
|
||||||
|
|
||||||
|
if is_prefill_phase:
|
||||||
|
self.dllm_phase = DllmReqPhase.STAGING_PREFILL
|
||||||
|
else:
|
||||||
|
self.dllm_phase = DllmReqPhase.STAGING_DECODE
|
||||||
|
|
||||||
|
def _init_fill_ids_for_dllm(self: Req):
|
||||||
|
if not self.dllm_ids:
|
||||||
|
self.dllm_ids = (
|
||||||
|
self.origin_input_ids
|
||||||
|
+ [self.dllm_config.mask_id] * self.dllm_config.block_size
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.dllm_block_offset += self.dllm_config.block_size
|
||||||
|
self.dllm_ids += [self.dllm_config.mask_id] * self.dllm_config.block_size
|
||||||
|
|
||||||
|
self.fill_ids = self.dllm_ids
|
||||||
@@ -0,0 +1,313 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import time
|
||||||
|
from typing import TYPE_CHECKING, List, Optional, Set, Union
|
||||||
|
|
||||||
|
from sglang.srt.dllm.config import DllmConfig
|
||||||
|
from sglang.srt.dllm.mixin.req import DllmReqPhase
|
||||||
|
from sglang.srt.managers.schedule_batch import Req, RequestStage, ScheduleBatch
|
||||||
|
from sglang.srt.managers.schedule_policy import AddReqResult, PrefillAdder
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.managers.scheduler import Scheduler
|
||||||
|
|
||||||
|
|
||||||
|
class SchedulerDllmMixin:
|
||||||
|
def init_diffusion_llm(self: Scheduler):
|
||||||
|
self.dllm_config = (
|
||||||
|
DllmConfig.from_server_args(self.server_args)
|
||||||
|
if self.server_args.dllm_algorithm is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
self.dllm_manager = DllmManager(dllm_config=self.dllm_config)
|
||||||
|
|
||||||
|
def get_new_batch_dllm(self: Scheduler) -> Optional[ScheduleBatch]:
|
||||||
|
"""Generate a new batch for DLLM (Diffusion LLM) scheduling."""
|
||||||
|
if self.try_preemption:
|
||||||
|
self.running_batch.batch_is_full = False
|
||||||
|
|
||||||
|
# Early exit if batch is full or no requests available
|
||||||
|
if self._should_skip_prefill():
|
||||||
|
return None
|
||||||
|
|
||||||
|
running_bs = len(self.running_batch.reqs)
|
||||||
|
self.policy.calc_priority(self.waiting_queue)
|
||||||
|
|
||||||
|
# Create prefill adder with resource constraints
|
||||||
|
adder = self._create_dllm_prefill_adder(running_bs)
|
||||||
|
|
||||||
|
# Initialize DLLM manager and transfer requests
|
||||||
|
self.dllm_manager.init_next_round()
|
||||||
|
self._fetch_waiting_reqs()
|
||||||
|
|
||||||
|
# Process batches
|
||||||
|
forward_mode = self._process_dllm_batches(adder)
|
||||||
|
|
||||||
|
can_run_list = adder.can_run_list
|
||||||
|
if not can_run_list:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Record metrics and update state
|
||||||
|
self._update_metrics_and_state_for_batch(can_run_list, adder, running_bs)
|
||||||
|
|
||||||
|
# Create and prepare batch
|
||||||
|
new_batch = self._create_dllm_batch(can_run_list, forward_mode)
|
||||||
|
return new_batch
|
||||||
|
|
||||||
|
def _fetch_waiting_reqs(self: Scheduler):
|
||||||
|
# Calculate how many requests can be added to DLLM manager
|
||||||
|
max_dllm_capacity = self.server_args.max_running_requests - len(
|
||||||
|
self.dllm_manager.waiting_queue
|
||||||
|
)
|
||||||
|
num_requests_to_add = min(max_dllm_capacity, len(self.waiting_queue))
|
||||||
|
|
||||||
|
if num_requests_to_add > 0:
|
||||||
|
requests_to_add = self.waiting_queue[:num_requests_to_add]
|
||||||
|
self.dllm_manager.add_waiting_reqs(requests_to_add)
|
||||||
|
self.waiting_queue = self.waiting_queue[num_requests_to_add:]
|
||||||
|
|
||||||
|
def _should_skip_prefill(self: Scheduler) -> bool:
|
||||||
|
"""Check if DLLM prefill should be skipped."""
|
||||||
|
if (
|
||||||
|
self.running_batch.batch_is_full or not self.waiting_queue
|
||||||
|
) and self.dllm_manager.is_empty():
|
||||||
|
return True
|
||||||
|
|
||||||
|
running_bs = len(self.running_batch.reqs)
|
||||||
|
if (
|
||||||
|
self.get_num_allocatable_reqs(running_bs) <= 0
|
||||||
|
and self.dllm_manager.is_empty()
|
||||||
|
and not self.try_preemption
|
||||||
|
):
|
||||||
|
self.running_batch.batch_is_full = True
|
||||||
|
return True
|
||||||
|
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _create_dllm_prefill_adder(self: Scheduler, running_bs: int) -> PrefillAdder:
|
||||||
|
"""Create a prefill adder configured for DLLM scheduling."""
|
||||||
|
return PrefillAdder(
|
||||||
|
self.page_size,
|
||||||
|
self.tree_cache,
|
||||||
|
self.token_to_kv_pool_allocator,
|
||||||
|
self.running_batch,
|
||||||
|
self.new_token_ratio,
|
||||||
|
self.max_prefill_tokens,
|
||||||
|
self.chunked_prefill_size,
|
||||||
|
running_bs if self.is_mixed_chunk else 0,
|
||||||
|
self.priority_scheduling_preemption_threshold,
|
||||||
|
prefill_max_requests=self.server_args.prefill_max_requests,
|
||||||
|
dllm_config=self.dllm_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _process_dllm_batches(self: Scheduler, adder: PrefillAdder) -> ForwardMode:
|
||||||
|
"""Process prefill or decode batches for DLLM."""
|
||||||
|
forward_mode = ForwardMode.DLLM_EXTEND
|
||||||
|
|
||||||
|
# Try prefill batch first
|
||||||
|
prefill_reqs = self.dllm_manager.get_prefill_requests()
|
||||||
|
if prefill_reqs:
|
||||||
|
self._process_batch_by_phase(
|
||||||
|
adder,
|
||||||
|
prefill_reqs,
|
||||||
|
DllmReqPhase.STAGING_PREFILL,
|
||||||
|
DllmReqPhase.INCOMING_PREFILL,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Fall back to decode batch
|
||||||
|
decode_reqs = self.dllm_manager.get_decode_requests()
|
||||||
|
self._process_batch_by_phase(
|
||||||
|
adder,
|
||||||
|
decode_reqs,
|
||||||
|
DllmReqPhase.STAGING_DECODE,
|
||||||
|
DllmReqPhase.INCOMING_DECODE,
|
||||||
|
)
|
||||||
|
|
||||||
|
return forward_mode
|
||||||
|
|
||||||
|
def _process_batch_by_phase(
|
||||||
|
self,
|
||||||
|
adder: PrefillAdder,
|
||||||
|
batch: List[Req],
|
||||||
|
staging_phase: DllmReqPhase,
|
||||||
|
incoming_phase: DllmReqPhase,
|
||||||
|
) -> None:
|
||||||
|
"""Process a batch, separating staging and incoming requests."""
|
||||||
|
staging_reqs = [req for req in batch if req.dllm_phase == staging_phase]
|
||||||
|
if staging_reqs:
|
||||||
|
staging_result = self.process_dllm_staging_reqs(adder, staging_reqs)
|
||||||
|
if staging_result != AddReqResult.CONTINUE:
|
||||||
|
return
|
||||||
|
|
||||||
|
incoming_reqs = [req for req in batch if req.dllm_phase == incoming_phase]
|
||||||
|
if incoming_reqs:
|
||||||
|
self.process_dllm_incoming_reqs(adder, incoming_reqs)
|
||||||
|
|
||||||
|
def _update_metrics_and_state_for_batch(
|
||||||
|
self: Scheduler, can_run_list: List[Req], adder: PrefillAdder, running_bs: int
|
||||||
|
) -> None:
|
||||||
|
"""Update metrics and state for the batch."""
|
||||||
|
if self.enable_metrics:
|
||||||
|
for req in can_run_list:
|
||||||
|
req.add_latency(RequestStage.PREFILL_WAITING)
|
||||||
|
|
||||||
|
if adder.preempt_list:
|
||||||
|
for req in adder.preempt_list:
|
||||||
|
self._add_request_to_queue(req)
|
||||||
|
|
||||||
|
if can_run_list:
|
||||||
|
self.dllm_manager.add_staging_reqs(can_run_list)
|
||||||
|
self.dllm_manager.increment_chunked_count()
|
||||||
|
|
||||||
|
self.adder = adder
|
||||||
|
self.can_run_list = can_run_list
|
||||||
|
self.running_bs = len(self.running_batch.reqs)
|
||||||
|
|
||||||
|
for req in can_run_list:
|
||||||
|
if req.time_stats.forward_entry_time == 0:
|
||||||
|
req.time_stats.forward_entry_time = time.perf_counter()
|
||||||
|
if self.enable_metrics:
|
||||||
|
self.metrics_collector.observe_queue_time(
|
||||||
|
req.time_stats.get_queueing_time(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _create_dllm_batch(
|
||||||
|
self: Scheduler, can_run_list: List[Req], forward_mode: ForwardMode
|
||||||
|
) -> ScheduleBatch:
|
||||||
|
"""Create and prepare a new DLLM batch."""
|
||||||
|
new_batch = ScheduleBatch.init_new(
|
||||||
|
can_run_list,
|
||||||
|
self.req_to_token_pool,
|
||||||
|
self.token_to_kv_pool_allocator,
|
||||||
|
self.tree_cache,
|
||||||
|
self.model_config,
|
||||||
|
self.enable_overlap,
|
||||||
|
self.spec_algorithm,
|
||||||
|
dllm_config=self.dllm_config,
|
||||||
|
)
|
||||||
|
new_batch.prepare_for_extend()
|
||||||
|
new_batch.forward_mode = forward_mode
|
||||||
|
new_batch.decoding_reqs = None
|
||||||
|
return new_batch
|
||||||
|
|
||||||
|
def process_dllm_incoming_reqs(
|
||||||
|
self: Scheduler, adder: PrefillAdder, reqs: List[Req]
|
||||||
|
) -> AddReqResult:
|
||||||
|
"""Process incoming DLLM requests with resource allocation and preemption."""
|
||||||
|
res = AddReqResult.CONTINUE
|
||||||
|
for req in reqs:
|
||||||
|
# Check if batch is full
|
||||||
|
running_bs = len(self.running_batch.reqs)
|
||||||
|
if len(adder.can_run_list) >= self.get_num_allocatable_reqs(running_bs):
|
||||||
|
self.running_batch.batch_is_full = True
|
||||||
|
|
||||||
|
# Try preemption if batch is full
|
||||||
|
if self.running_batch.batch_is_full:
|
||||||
|
if not self.try_preemption or not adder.preempt_to_schedule(
|
||||||
|
req, self.server_args
|
||||||
|
):
|
||||||
|
break
|
||||||
|
|
||||||
|
# Prepare and add request
|
||||||
|
req.init_next_round_input(self.tree_cache)
|
||||||
|
res = adder.add_one_req(
|
||||||
|
req,
|
||||||
|
has_chunked_req=True,
|
||||||
|
truncation_align_size=self.truncation_align_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
if res != AddReqResult.CONTINUE:
|
||||||
|
if res == AddReqResult.NO_TOKEN:
|
||||||
|
self.running_batch.batch_is_full = True
|
||||||
|
break
|
||||||
|
|
||||||
|
return res
|
||||||
|
|
||||||
|
def process_dllm_staging_reqs(
|
||||||
|
self: Scheduler, adder: PrefillAdder, reqs: List[Req]
|
||||||
|
) -> AddReqResult:
|
||||||
|
"""Process staging DLLM requests with resource allocation."""
|
||||||
|
for req in reqs:
|
||||||
|
res = adder.add_dllm_staging_req(req)
|
||||||
|
if res == AddReqResult.NO_TOKEN:
|
||||||
|
return res
|
||||||
|
|
||||||
|
return AddReqResult.CONTINUE
|
||||||
|
|
||||||
|
|
||||||
|
class DllmManager:
|
||||||
|
"""
|
||||||
|
Manager for Diffusion LLM request scheduling.
|
||||||
|
|
||||||
|
Maintains two queues:
|
||||||
|
- waiting_queue: The requests waiting to be scheduled with max running requests limit
|
||||||
|
- staging_queue: Requests allocated resources by PrefillAdder
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, dllm_config: Optional[DllmConfig] = None):
|
||||||
|
self.dllm_config = dllm_config
|
||||||
|
self.max_running_reqs = (
|
||||||
|
dllm_config.max_running_requests if dllm_config is not None else 1
|
||||||
|
)
|
||||||
|
self.waiting_queue: List[Req] = []
|
||||||
|
self.staging_queue: List[Req] = []
|
||||||
|
|
||||||
|
def get_prefill_requests(self) -> List[Req]:
|
||||||
|
"""Get all prefill requests from waiting queue."""
|
||||||
|
return [req for req in self.waiting_queue if req.is_dllm_prefill()]
|
||||||
|
|
||||||
|
def get_decode_requests(self) -> List[Req]:
|
||||||
|
"""Get all decode requests from waiting queue."""
|
||||||
|
return [req for req in self.waiting_queue if not req.is_dllm_prefill()]
|
||||||
|
|
||||||
|
def add_waiting_reqs(self, reqs: Union[Req, List[Req]]) -> None:
|
||||||
|
"""Add requests to waiting queue with redundancy check."""
|
||||||
|
assert self.dllm_config is not None, "Diffusion LLM config is not set."
|
||||||
|
|
||||||
|
reqs_to_add = reqs if isinstance(reqs, list) else [reqs]
|
||||||
|
|
||||||
|
# Check for duplicate request IDs
|
||||||
|
if self._has_duplicate_reqs(reqs_to_add):
|
||||||
|
raise RuntimeError("Redundant requests detected in dLLM requests.")
|
||||||
|
|
||||||
|
self.waiting_queue.extend(reqs_to_add)
|
||||||
|
|
||||||
|
def add_staging_reqs(self, reqs: Union[Req, List[Req]]) -> None:
|
||||||
|
"""Add requests to staging queue (allocated by PrefillAdder)."""
|
||||||
|
reqs_to_add = reqs if isinstance(reqs, list) else [reqs]
|
||||||
|
self.staging_queue.extend(reqs_to_add)
|
||||||
|
|
||||||
|
def _has_duplicate_reqs(self, reqs: List[Req]) -> bool:
|
||||||
|
"""Check if any request ID already exists in waiting queue."""
|
||||||
|
existing_rids: Set[str] = {r.rid for r in self.waiting_queue}
|
||||||
|
return any(req.rid in existing_rids for req in reqs)
|
||||||
|
|
||||||
|
def any_staging_reqs(self) -> bool:
|
||||||
|
"""Check if there are requests in staging queue."""
|
||||||
|
return self.dllm_config is not None and len(self.staging_queue) > 0
|
||||||
|
|
||||||
|
def is_empty(self) -> bool:
|
||||||
|
"""Check if both queues are empty or DLLM is not configured."""
|
||||||
|
if self.dllm_config is None:
|
||||||
|
return True
|
||||||
|
return len(self.waiting_queue) == 0
|
||||||
|
|
||||||
|
def increment_chunked_count(self) -> None:
|
||||||
|
"""Increment chunked count for all staging requests."""
|
||||||
|
for req in self.staging_queue:
|
||||||
|
req.is_chunked += 1
|
||||||
|
|
||||||
|
def filter_finished_reqs(self) -> None:
|
||||||
|
"""Remove finished requests from both queues."""
|
||||||
|
self.waiting_queue = [req for req in self.waiting_queue if not req.finished()]
|
||||||
|
self.staging_queue = [req for req in self.staging_queue if not req.finished()]
|
||||||
|
|
||||||
|
def init_next_round(self) -> None:
|
||||||
|
"""Initialize staging requests for next round and clear staging queue."""
|
||||||
|
for req in self.staging_queue:
|
||||||
|
req.init_next_round_input()
|
||||||
|
self.staging_queue = []
|
||||||
@@ -58,6 +58,7 @@ from sglang.srt.disaggregation.decode_schedule_batch_mixin import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank
|
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank
|
||||||
|
from sglang.srt.dllm.mixin.req import ReqDllmMixin
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
|
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
|
||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
@@ -507,7 +508,7 @@ class RequestStage(str, enum.Enum):
|
|||||||
DECODE_QUICK_FINISH = "quick_finish"
|
DECODE_QUICK_FINISH = "quick_finish"
|
||||||
|
|
||||||
|
|
||||||
class Req:
|
class Req(ReqDllmMixin):
|
||||||
"""The input and output status of a request."""
|
"""The input and output status of a request."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -804,9 +805,7 @@ class Req:
|
|||||||
self.dimensions = dimensions
|
self.dimensions = dimensions
|
||||||
|
|
||||||
# For diffusion LLM
|
# For diffusion LLM
|
||||||
self.dllm_ids = []
|
self.init_diffusion_llm(dllm_config)
|
||||||
self.dllm_block_offset = 0
|
|
||||||
self.dllm_config = dllm_config
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def seqlen(self) -> int:
|
def seqlen(self) -> int:
|
||||||
@@ -868,23 +867,10 @@ class Req:
|
|||||||
# Whether request reached finished condition
|
# Whether request reached finished condition
|
||||||
return self.finished_reason is not None
|
return self.finished_reason is not None
|
||||||
|
|
||||||
def is_dllm(self):
|
|
||||||
return self.dllm_config is not None
|
|
||||||
|
|
||||||
def _init_fill_ids_for_dllm(self):
|
|
||||||
if not self.dllm_ids:
|
|
||||||
self.dllm_ids = (
|
|
||||||
self.origin_input_ids
|
|
||||||
+ [self.dllm_config.mask_id] * self.dllm_config.block_size
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
self.dllm_block_offset += self.dllm_config.block_size
|
|
||||||
self.dllm_ids += [self.dllm_config.mask_id] * self.dllm_config.block_size
|
|
||||||
self.fill_ids = self.dllm_ids
|
|
||||||
|
|
||||||
def init_next_round_input(self, tree_cache: Optional[BasePrefixCache] = None):
|
def init_next_round_input(self, tree_cache: Optional[BasePrefixCache] = None):
|
||||||
if self.is_dllm():
|
if self.is_dllm():
|
||||||
self._init_fill_ids_for_dllm()
|
self._init_fill_ids_for_dllm()
|
||||||
|
self.determine_dllm_phase()
|
||||||
else:
|
else:
|
||||||
self.fill_ids = self.origin_input_ids + self.output_ids
|
self.fill_ids = self.origin_input_ids + self.output_ids
|
||||||
|
|
||||||
@@ -1194,62 +1180,6 @@ class Req:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class DllmStagingReqs:
|
|
||||||
def __init__(self, dllm_config: Optional[DllmConfig] = None):
|
|
||||||
self.dllm_config = dllm_config
|
|
||||||
self.max_running_reqs = (
|
|
||||||
dllm_config.max_running_requests if dllm_config is not None else 1
|
|
||||||
)
|
|
||||||
self.reqs: List[Req] = []
|
|
||||||
|
|
||||||
def add_reqs(self, req: Union[Req, List[Req], "DllmStagingReqs"]):
|
|
||||||
assert self.dllm_config is not None, "Diffusion LLM config is not set."
|
|
||||||
|
|
||||||
if isinstance(req, DllmStagingReqs):
|
|
||||||
reqs_to_add = req.reqs
|
|
||||||
elif isinstance(req, list):
|
|
||||||
reqs_to_add = req
|
|
||||||
else:
|
|
||||||
reqs_to_add = [req]
|
|
||||||
|
|
||||||
num_to_add = len(reqs_to_add)
|
|
||||||
|
|
||||||
# Sanity check:
|
|
||||||
if self.check_redundant_reqs(reqs_to_add):
|
|
||||||
raise RuntimeError("Redundant requests detected in dLLM requests.")
|
|
||||||
|
|
||||||
if len(self.reqs) + num_to_add > self.max_running_reqs:
|
|
||||||
raise RuntimeError(
|
|
||||||
f"Exceeding maximum number of concurrent diffusion LLM requests: {self.max_running_reqs}"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.reqs.extend(reqs_to_add)
|
|
||||||
|
|
||||||
def check_redundant_reqs(self, reqs: List[Req]) -> bool:
|
|
||||||
existing_rids: Set[str] = {r.rid for r in self.reqs}
|
|
||||||
return any(req.rid in existing_rids for req in reqs)
|
|
||||||
|
|
||||||
def init_next_round(self):
|
|
||||||
for req in self.reqs:
|
|
||||||
req.init_next_round_input()
|
|
||||||
|
|
||||||
def non_empty(self) -> bool:
|
|
||||||
return self.dllm_config is not None and len(self.reqs) > 0
|
|
||||||
|
|
||||||
def empty(self) -> bool:
|
|
||||||
return self.dllm_config is None or len(self.reqs) == 0
|
|
||||||
|
|
||||||
def update_chunked_status(self):
|
|
||||||
for req in self.reqs:
|
|
||||||
req.is_chunked += 1
|
|
||||||
|
|
||||||
def filter_finished_reqs(self):
|
|
||||||
self.reqs = [req for req in self.reqs if not req.finished()]
|
|
||||||
|
|
||||||
def __iter__(self):
|
|
||||||
return iter(self.reqs)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||||
"""Store all information of a batch on the scheduler."""
|
"""Store all information of a batch on the scheduler."""
|
||||||
@@ -1370,7 +1300,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
hicache_consumer_index: int = -1
|
hicache_consumer_index: int = -1
|
||||||
|
|
||||||
# Diffusion LLM
|
# Diffusion LLM
|
||||||
dllm_staging_reqs: Optional[DllmStagingReqs] = None
|
|
||||||
dllm_config: Optional[DllmConfig] = None
|
dllm_config: Optional[DllmConfig] = None
|
||||||
|
|
||||||
# Metrics
|
# Metrics
|
||||||
@@ -1387,7 +1316,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
enable_overlap: bool,
|
enable_overlap: bool,
|
||||||
spec_algorithm: SpeculativeAlgorithm,
|
spec_algorithm: SpeculativeAlgorithm,
|
||||||
chunked_req: Optional[Req] = None,
|
chunked_req: Optional[Req] = None,
|
||||||
dllm_staging_reqs: Optional[DllmStagingReqs] = None,
|
|
||||||
dllm_config: Optional[DllmConfig] = None,
|
dllm_config: Optional[DllmConfig] = None,
|
||||||
):
|
):
|
||||||
return_logprob = any(req.return_logprob for req in reqs)
|
return_logprob = any(req.return_logprob for req in reqs)
|
||||||
@@ -1413,7 +1341,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
return_routed_experts=any(req.return_routed_experts for req in reqs),
|
return_routed_experts=any(req.return_routed_experts for req in reqs),
|
||||||
is_prefill_only=all(req.is_prefill_only for req in reqs),
|
is_prefill_only=all(req.is_prefill_only for req in reqs),
|
||||||
chunked_req=chunked_req,
|
chunked_req=chunked_req,
|
||||||
dllm_staging_reqs=dllm_staging_reqs,
|
|
||||||
dllm_config=dllm_config,
|
dllm_config=dllm_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.dllm.config import DllmConfig
|
from sglang.srt.dllm.config import DllmConfig
|
||||||
from sglang.srt.layers.attention.nsa.utils import is_nsa_prefill_cp_in_seq_split
|
from sglang.srt.layers.attention.nsa.utils import is_nsa_prefill_cp_in_seq_split
|
||||||
from sglang.srt.managers.schedule_batch import DllmStagingReqs, Req, ScheduleBatch
|
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
BasePrefixCache,
|
BasePrefixCache,
|
||||||
InsertParams,
|
InsertParams,
|
||||||
@@ -435,7 +435,6 @@ class PrefillAdder:
|
|||||||
max_running_reqs = dllm_config.max_running_requests
|
max_running_reqs = dllm_config.max_running_requests
|
||||||
|
|
||||||
self.rem_dllm_tokens = max_running_reqs * self.dllm_block_size
|
self.rem_dllm_tokens = max_running_reqs * self.dllm_block_size
|
||||||
self.dllm_staging_reqs = DllmStagingReqs(dllm_config=dllm_config)
|
|
||||||
|
|
||||||
def _get_running_request_total_token_offset(self, req: Req) -> int:
|
def _get_running_request_total_token_offset(self, req: Req) -> int:
|
||||||
return (
|
return (
|
||||||
@@ -551,7 +550,6 @@ class PrefillAdder:
|
|||||||
req.fill_ids = req.fill_ids[: prefix_len + trunc_len]
|
req.fill_ids = req.fill_ids[: prefix_len + trunc_len]
|
||||||
|
|
||||||
self.can_run_list.append(req)
|
self.can_run_list.append(req)
|
||||||
self.dllm_staging_reqs.add_reqs(req)
|
|
||||||
|
|
||||||
self._update_prefill_budget(prefix_len, trunc_len, 0)
|
self._update_prefill_budget(prefix_len, trunc_len, 0)
|
||||||
|
|
||||||
@@ -562,6 +560,34 @@ class PrefillAdder:
|
|||||||
else:
|
else:
|
||||||
self.tree_cache.inc_lock_ref(req.last_node)
|
self.tree_cache.inc_lock_ref(req.last_node)
|
||||||
|
|
||||||
|
def add_dllm_staging_req(self, req: Req):
|
||||||
|
assert self.dllm_config is not None
|
||||||
|
_rem_tokens = self._get_dllm_remain_tokens()
|
||||||
|
|
||||||
|
if _rem_tokens <= 0:
|
||||||
|
return AddReqResult.NO_TOKEN
|
||||||
|
|
||||||
|
# Truncate input length to available tokens and update request metadata
|
||||||
|
truncated = req.extend_input_len > _rem_tokens
|
||||||
|
req.extend_input_len = min(req.extend_input_len, _rem_tokens)
|
||||||
|
req.fill_ids = req.fill_ids[: len(req.prefix_indices) + req.extend_input_len]
|
||||||
|
self.can_run_list.append(req)
|
||||||
|
|
||||||
|
# Update budget: reserve max_new_tokens only if not truncated
|
||||||
|
max_new_tokens = (
|
||||||
|
min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS)
|
||||||
|
if not truncated
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
self._update_prefill_budget(0, req.extend_input_len, max_new_tokens)
|
||||||
|
|
||||||
|
# Return based on remaining token availability
|
||||||
|
return (
|
||||||
|
AddReqResult.NO_TOKEN
|
||||||
|
if self._get_dllm_remain_tokens() <= 0
|
||||||
|
else AddReqResult.CONTINUE
|
||||||
|
)
|
||||||
|
|
||||||
def add_chunked_req(self, req: Req):
|
def add_chunked_req(self, req: Req):
|
||||||
if self.dllm_config is not None:
|
if self.dllm_config is not None:
|
||||||
_rem_tokens = self._get_dllm_remain_tokens()
|
_rem_tokens = self._get_dllm_remain_tokens()
|
||||||
|
|||||||
@@ -58,7 +58,7 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.distributed import get_pp_group, get_world_group
|
from sglang.srt.distributed import get_pp_group, get_world_group
|
||||||
from sglang.srt.distributed.parallel_state import get_tp_group
|
from sglang.srt.distributed.parallel_state import get_tp_group
|
||||||
from sglang.srt.dllm.config import DllmConfig
|
from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
@@ -144,7 +144,6 @@ from sglang.srt.managers.schedule_batch import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_policy import (
|
from sglang.srt.managers.schedule_policy import (
|
||||||
AddReqResult,
|
AddReqResult,
|
||||||
DllmStagingReqs,
|
|
||||||
PrefillAdder,
|
PrefillAdder,
|
||||||
SchedulePolicy,
|
SchedulePolicy,
|
||||||
)
|
)
|
||||||
@@ -255,6 +254,7 @@ class Scheduler(
|
|||||||
SchedulerRuntimeCheckerMixin,
|
SchedulerRuntimeCheckerMixin,
|
||||||
SchedulerPPMixin,
|
SchedulerPPMixin,
|
||||||
SchedulerDPAttnMixin,
|
SchedulerDPAttnMixin,
|
||||||
|
SchedulerDllmMixin,
|
||||||
):
|
):
|
||||||
"""A scheduler that manages a tensor parallel GPU worker."""
|
"""A scheduler that manages a tensor parallel GPU worker."""
|
||||||
|
|
||||||
@@ -399,11 +399,6 @@ class Scheduler(
|
|||||||
|
|
||||||
def init_model_config(self):
|
def init_model_config(self):
|
||||||
self.model_config = ModelConfig.from_server_args(self.server_args)
|
self.model_config = ModelConfig.from_server_args(self.server_args)
|
||||||
self.dllm_config = ( # For diffusion LLM
|
|
||||||
DllmConfig.from_server_args(self.server_args)
|
|
||||||
if self.server_args.dllm_algorithm is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
|
|
||||||
def init_ipc_channels(self, port_args: PortArgs):
|
def init_ipc_channels(self, port_args: PortArgs):
|
||||||
context = zmq.Context(2)
|
context = zmq.Context(2)
|
||||||
@@ -763,9 +758,6 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
self.enable_dynamic_chunking = False
|
self.enable_dynamic_chunking = False
|
||||||
|
|
||||||
def init_diffusion_llm(self):
|
|
||||||
self.dllm_staging_reqs = DllmStagingReqs(dllm_config=self.dllm_config)
|
|
||||||
|
|
||||||
def init_schedule_policy(self):
|
def init_schedule_policy(self):
|
||||||
# Init schedule policy and new token estimation
|
# Init schedule policy and new token estimation
|
||||||
self.policy = SchedulePolicy(
|
self.policy = SchedulePolicy(
|
||||||
@@ -1323,7 +1315,7 @@ class Scheduler(
|
|||||||
# If it is a health check generation request and there are running requests, ignore it.
|
# If it is a health check generation request and there are running requests, ignore it.
|
||||||
if is_health_check_generate_req(recv_req) and (
|
if is_health_check_generate_req(recv_req) and (
|
||||||
self.chunked_req is not None
|
self.chunked_req is not None
|
||||||
or self.dllm_staging_reqs.non_empty()
|
or self.dllm_manager.any_staging_reqs()
|
||||||
or not self.running_batch.is_empty()
|
or not self.running_batch.is_empty()
|
||||||
or len(self.offload_tags) > 0
|
or len(self.offload_tags) > 0
|
||||||
):
|
):
|
||||||
@@ -1830,20 +1822,15 @@ class Scheduler(
|
|||||||
def get_next_batch_to_run(self) -> Optional[ScheduleBatch]:
|
def get_next_batch_to_run(self) -> Optional[ScheduleBatch]:
|
||||||
self._abort_on_queued_timeout()
|
self._abort_on_queued_timeout()
|
||||||
if self.dllm_config is not None:
|
if self.dllm_config is not None:
|
||||||
self.dllm_staging_reqs.filter_finished_reqs()
|
self.dllm_manager.filter_finished_reqs()
|
||||||
|
|
||||||
# Merge the prefill batch into the running batch
|
# Merge the prefill batch into the running batch
|
||||||
chunked_req_to_exclude = set()
|
chunked_req_to_exclude = set()
|
||||||
|
|
||||||
if self.dllm_config is not None:
|
if self.dllm_config is not None and self.dllm_manager.any_staging_reqs():
|
||||||
assert (
|
chunked_req_to_exclude.update(self.dllm_manager.staging_queue)
|
||||||
self.chunked_req is None
|
for req in self.dllm_manager.staging_queue:
|
||||||
), "chunked_req should be None when dllm_config is set"
|
self.stash_chunked_request(req)
|
||||||
|
|
||||||
if self.dllm_staging_reqs.non_empty():
|
|
||||||
chunked_req_to_exclude.update(self.dllm_staging_reqs)
|
|
||||||
for req in self.dllm_staging_reqs:
|
|
||||||
self.stash_chunked_request(req)
|
|
||||||
|
|
||||||
if self.chunked_req is not None:
|
if self.chunked_req is not None:
|
||||||
# Move the chunked request out of the batch so that we can merge
|
# Move the chunked request out of the batch so that we can merge
|
||||||
@@ -1857,8 +1844,8 @@ class Scheduler(
|
|||||||
# We need to discard it.
|
# We need to discard it.
|
||||||
chunked_req_to_exclude.add(self.last_batch.chunked_req)
|
chunked_req_to_exclude.add(self.last_batch.chunked_req)
|
||||||
|
|
||||||
if self.last_batch.dllm_staging_reqs.non_empty():
|
if self.dllm_config is not None and self.last_batch.reqs:
|
||||||
chunked_req_to_exclude.update(self.last_batch.dllm_staging_reqs)
|
chunked_req_to_exclude.update(self.last_batch.reqs)
|
||||||
|
|
||||||
# Filter batch
|
# Filter batch
|
||||||
last_bs = self.last_batch.batch_size()
|
last_bs = self.last_batch.batch_size()
|
||||||
@@ -1877,7 +1864,10 @@ class Scheduler(
|
|||||||
# Merge running_batch with prefill batch
|
# Merge running_batch with prefill batch
|
||||||
self.running_batch.merge_batch(self.last_batch)
|
self.running_batch.merge_batch(self.last_batch)
|
||||||
|
|
||||||
new_batch = self.get_new_batch_prefill()
|
if self.dllm_config is not None:
|
||||||
|
new_batch = self.get_new_batch_dllm()
|
||||||
|
else:
|
||||||
|
new_batch = self.get_new_batch_prefill()
|
||||||
|
|
||||||
need_mlp_sync = self.require_mlp_sync
|
need_mlp_sync = self.require_mlp_sync
|
||||||
if need_mlp_sync and not self.spec_algorithm.is_none():
|
if need_mlp_sync and not self.spec_algorithm.is_none():
|
||||||
@@ -1947,9 +1937,9 @@ class Scheduler(
|
|||||||
# Reset batch_is_full to try preemption with a prefill adder.
|
# Reset batch_is_full to try preemption with a prefill adder.
|
||||||
self.running_batch.batch_is_full = False
|
self.running_batch.batch_is_full = False
|
||||||
|
|
||||||
if (self.running_batch.batch_is_full or len(self.waiting_queue) == 0) and (
|
if (
|
||||||
not self.dllm_staging_reqs.non_empty() and self.chunked_req is None
|
self.running_batch.batch_is_full or len(self.waiting_queue) == 0
|
||||||
):
|
) and self.chunked_req is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
running_bs = len(self.running_batch.reqs)
|
running_bs = len(self.running_batch.reqs)
|
||||||
@@ -1960,7 +1950,7 @@ class Scheduler(
|
|||||||
# Instead, we should always allow chunked requests to be added, otherwise, there will be a memory leak.
|
# Instead, we should always allow chunked requests to be added, otherwise, there will be a memory leak.
|
||||||
if (
|
if (
|
||||||
self.get_num_allocatable_reqs(running_bs) <= 0
|
self.get_num_allocatable_reqs(running_bs) <= 0
|
||||||
and (self.dllm_staging_reqs.empty() or self.chunked_req is not None)
|
and self.chunked_req is not None
|
||||||
and not self.try_preemption
|
and not self.try_preemption
|
||||||
):
|
):
|
||||||
self.running_batch.batch_is_full = True
|
self.running_batch.batch_is_full = True
|
||||||
@@ -2002,17 +1992,6 @@ class Scheduler(
|
|||||||
dllm_config=self.dllm_config,
|
dllm_config=self.dllm_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.dllm_config is not None:
|
|
||||||
assert (
|
|
||||||
self.chunked_req is None
|
|
||||||
), "chunked_req should be None when dllm_config is set"
|
|
||||||
|
|
||||||
if self.dllm_staging_reqs.non_empty():
|
|
||||||
self.dllm_staging_reqs.init_next_round()
|
|
||||||
for req in self.dllm_staging_reqs:
|
|
||||||
adder.add_chunked_req(req)
|
|
||||||
self.dllm_staging_reqs.add_reqs(adder.dllm_staging_reqs)
|
|
||||||
|
|
||||||
if self.chunked_req is not None:
|
if self.chunked_req is not None:
|
||||||
self.chunked_req.init_next_round_input()
|
self.chunked_req.init_next_round_input()
|
||||||
self.chunked_req = adder.add_chunked_req(self.chunked_req)
|
self.chunked_req = adder.add_chunked_req(self.chunked_req)
|
||||||
@@ -2066,9 +2045,7 @@ class Scheduler(
|
|||||||
req.init_next_round_input(self.tree_cache)
|
req.init_next_round_input(self.tree_cache)
|
||||||
res = adder.add_one_req(
|
res = adder.add_one_req(
|
||||||
req,
|
req,
|
||||||
has_chunked_req=(
|
has_chunked_req=(self.chunked_req is not None),
|
||||||
self.dllm_staging_reqs.non_empty() or self.chunked_req is not None
|
|
||||||
),
|
|
||||||
truncation_align_size=self.truncation_align_size,
|
truncation_align_size=self.truncation_align_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2103,14 +2080,6 @@ class Scheduler(
|
|||||||
for req in adder.preempt_list:
|
for req in adder.preempt_list:
|
||||||
self._add_request_to_queue(req)
|
self._add_request_to_queue(req)
|
||||||
|
|
||||||
if self.dllm_config is not None:
|
|
||||||
assert (
|
|
||||||
self.chunked_req is None
|
|
||||||
), "chunked_req should be None when dllm_config is set"
|
|
||||||
|
|
||||||
if adder.dllm_staging_reqs.non_empty():
|
|
||||||
self.dllm_staging_reqs.add_reqs(adder.dllm_staging_reqs)
|
|
||||||
|
|
||||||
if adder.new_chunked_req is not None:
|
if adder.new_chunked_req is not None:
|
||||||
# Update chunked prefill
|
# Update chunked prefill
|
||||||
assert self.chunked_req is None
|
assert self.chunked_req is None
|
||||||
@@ -2119,9 +2088,6 @@ class Scheduler(
|
|||||||
if self.chunked_req is not None:
|
if self.chunked_req is not None:
|
||||||
self.chunked_req.is_chunked += 1
|
self.chunked_req.is_chunked += 1
|
||||||
|
|
||||||
if self.dllm_staging_reqs.non_empty():
|
|
||||||
self.dllm_staging_reqs.update_chunked_status()
|
|
||||||
|
|
||||||
# Record for logging prefill stats after forward
|
# Record for logging prefill stats after forward
|
||||||
self.adder = adder
|
self.adder = adder
|
||||||
self.can_run_list = can_run_list
|
self.can_run_list = can_run_list
|
||||||
@@ -2146,8 +2112,6 @@ class Scheduler(
|
|||||||
self.enable_overlap,
|
self.enable_overlap,
|
||||||
self.spec_algorithm,
|
self.spec_algorithm,
|
||||||
chunked_req=self.chunked_req,
|
chunked_req=self.chunked_req,
|
||||||
dllm_staging_reqs=self.dllm_staging_reqs,
|
|
||||||
dllm_config=self.dllm_config,
|
|
||||||
)
|
)
|
||||||
if self.enable_hierarchical_cache:
|
if self.enable_hierarchical_cache:
|
||||||
# todo (zhiqiang): disable cuda graph execution if hicache loading triggered
|
# todo (zhiqiang): disable cuda graph execution if hicache loading triggered
|
||||||
|
|||||||
@@ -95,7 +95,7 @@ class ForwardMode(IntEnum):
|
|||||||
# Split Prefill for PD multiplexing
|
# Split Prefill for PD multiplexing
|
||||||
SPLIT_PREFILL = auto()
|
SPLIT_PREFILL = auto()
|
||||||
|
|
||||||
# Used in diffusion LLM inference
|
# Used in dLLM
|
||||||
DLLM_EXTEND = auto()
|
DLLM_EXTEND = auto()
|
||||||
|
|
||||||
def is_prefill(self):
|
def is_prefill(self):
|
||||||
|
|||||||
@@ -2758,6 +2758,24 @@ class ServerArgs:
|
|||||||
)
|
)
|
||||||
self.pp_size = 1
|
self.pp_size = 1
|
||||||
|
|
||||||
|
if self.enable_lora:
|
||||||
|
logger.warning(
|
||||||
|
"Currently LoRA is not supported by diffusion LLM inference."
|
||||||
|
)
|
||||||
|
self.enable_lora = False
|
||||||
|
|
||||||
|
if self.disaggregation_mode != "null":
|
||||||
|
logger.warning(
|
||||||
|
"Currently disaggregation is not supported by diffusion LLM inference."
|
||||||
|
)
|
||||||
|
self.disaggregation_mode = "null"
|
||||||
|
|
||||||
|
if self.enable_mixed_chunk:
|
||||||
|
logger.warning(
|
||||||
|
"Mixed chunked prefill is disabled because of using diffusion LLM inference."
|
||||||
|
)
|
||||||
|
self.enable_mixed_chunk = False
|
||||||
|
|
||||||
def _handle_other_validations(self):
|
def _handle_other_validations(self):
|
||||||
# Handle model inference tensor dump.
|
# Handle model inference tensor dump.
|
||||||
if self.debug_tensor_dump_output_folder is not None:
|
if self.debug_tensor_dump_output_folder is not None:
|
||||||
|
|||||||
@@ -1,71 +0,0 @@
|
|||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
|
||||||
|
|
||||||
register_cuda_ci(est_time=500, suite="stage-b-test-large-1-gpu")
|
|
||||||
|
|
||||||
import unittest
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
from sglang.srt.utils import kill_process_tree
|
|
||||||
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
|
|
||||||
from sglang.test.test_utils import (
|
|
||||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
|
||||||
DEFAULT_URL_FOR_TEST,
|
|
||||||
CustomTestCase,
|
|
||||||
popen_launch_server,
|
|
||||||
)
|
|
||||||
|
|
||||||
"""
|
|
||||||
Test dLLM batching capability on CUDA GPUs.
|
|
||||||
|
|
||||||
As current dLLM batching performance is suboptimal to BS=1, this test only verifies correctness.
|
|
||||||
The test will be removed once dLLM batching performance improves.
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
class TestBatching(CustomTestCase):
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.model = "inclusionAI/LLaDA2.0-mini"
|
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
|
||||||
|
|
||||||
other_args = [
|
|
||||||
"--trust-remote-code",
|
|
||||||
"--mem-fraction-static",
|
|
||||||
"0.9",
|
|
||||||
"--max-running-requests",
|
|
||||||
"4",
|
|
||||||
"--attention-backend",
|
|
||||||
"flashinfer",
|
|
||||||
"--dllm-algorithm",
|
|
||||||
"LowConfidence",
|
|
||||||
]
|
|
||||||
|
|
||||||
cls.process = popen_launch_server(
|
|
||||||
cls.model,
|
|
||||||
cls.base_url,
|
|
||||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
|
||||||
other_args=other_args,
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
kill_process_tree(cls.process.pid)
|
|
||||||
|
|
||||||
def test_gsm8k(self):
|
|
||||||
args = SimpleNamespace(
|
|
||||||
num_shots=5,
|
|
||||||
data_path=None,
|
|
||||||
num_questions=200,
|
|
||||||
max_new_tokens=512,
|
|
||||||
parallel=128,
|
|
||||||
host="http://127.0.0.1",
|
|
||||||
port=int(self.base_url.split(":")[-1]),
|
|
||||||
)
|
|
||||||
metrics = run_eval_few_shot_gsm8k(args)
|
|
||||||
print(f"{metrics=}")
|
|
||||||
|
|
||||||
self.assertGreater(metrics["accuracy"], 0.88)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
unittest.main()
|
|
||||||
@@ -28,14 +28,21 @@ class TestLLaDA2Mini(CustomTestCase):
|
|||||||
|
|
||||||
other_args = [
|
other_args = [
|
||||||
"--trust-remote-code",
|
"--trust-remote-code",
|
||||||
|
"--tp-size",
|
||||||
|
"1",
|
||||||
"--mem-fraction-static",
|
"--mem-fraction-static",
|
||||||
"0.9",
|
"0.9",
|
||||||
"--max-running-requests",
|
"--max-running-requests",
|
||||||
"1",
|
"4",
|
||||||
"--attention-backend",
|
"--attention-backend",
|
||||||
"flashinfer",
|
"flashinfer",
|
||||||
"--dllm-algorithm",
|
"--dllm-algorithm",
|
||||||
"LowConfidence", # TODO: Add dLLM configurations
|
"LowConfidence",
|
||||||
|
"--cuda-graph-bs",
|
||||||
|
"1",
|
||||||
|
"2",
|
||||||
|
"3",
|
||||||
|
"4",
|
||||||
]
|
]
|
||||||
|
|
||||||
cls.process = popen_launch_server(
|
cls.process = popen_launch_server(
|
||||||
@@ -66,7 +73,7 @@ class TestLLaDA2Mini(CustomTestCase):
|
|||||||
if is_in_amd_ci():
|
if is_in_amd_ci():
|
||||||
self.assertGreater(metrics["output_throughput"], 80)
|
self.assertGreater(metrics["output_throughput"], 80)
|
||||||
else:
|
else:
|
||||||
self.assertGreater(metrics["output_throughput"], 150)
|
self.assertGreater(metrics["output_throughput"], 250)
|
||||||
|
|
||||||
def test_bs_1_speed(self):
|
def test_bs_1_speed(self):
|
||||||
args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048)
|
args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048)
|
||||||
|
|||||||
Reference in New Issue
Block a user