Pack scattered scheduler IPC channel state into a dedicated container (#25714)
This commit is contained in:
@@ -33,7 +33,6 @@ import psutil
|
|||||||
import setproctitle
|
import setproctitle
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed
|
import torch.distributed
|
||||||
import zmq
|
|
||||||
from torch.cuda import Stream as CudaStream
|
from torch.cuda import Stream as CudaStream
|
||||||
from torch.distributed import barrier
|
from torch.distributed import barrier
|
||||||
|
|
||||||
@@ -171,6 +170,9 @@ from sglang.srt.managers.scheduler_components.invariant_checker import (
|
|||||||
SchedulerInvariantChecker,
|
SchedulerInvariantChecker,
|
||||||
create_scheduler_watchdog,
|
create_scheduler_watchdog,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.scheduler_components.ipc_channels import (
|
||||||
|
SchedulerIpcChannels,
|
||||||
|
)
|
||||||
from sglang.srt.managers.scheduler_components.kv_events_publisher import (
|
from sglang.srt.managers.scheduler_components.kv_events_publisher import (
|
||||||
SchedulerKvEventsPublisher,
|
SchedulerKvEventsPublisher,
|
||||||
)
|
)
|
||||||
@@ -185,7 +187,6 @@ from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
|||||||
PrefillStats,
|
PrefillStats,
|
||||||
SchedulerMetricsReporter,
|
SchedulerMetricsReporter,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.scheduler_components.output_sender import SenderWrapper
|
|
||||||
from sglang.srt.managers.scheduler_components.output_streamer import (
|
from sglang.srt.managers.scheduler_components.output_streamer import (
|
||||||
SchedulerOutputStreamer,
|
SchedulerOutputStreamer,
|
||||||
)
|
)
|
||||||
@@ -250,7 +251,6 @@ from sglang.srt.utils.hf_transformers_utils import (
|
|||||||
get_tokenizer,
|
get_tokenizer,
|
||||||
get_tokenizer_from_processor,
|
get_tokenizer_from_processor,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.network import get_zmq_socket
|
|
||||||
from sglang.srt.utils.numa_utils import get_numa_node_if_available, numa_bind_to_node
|
from sglang.srt.utils.numa_utils import get_numa_node_if_available, numa_bind_to_node
|
||||||
from sglang.srt.utils.tensor_bridge import use_mlx
|
from sglang.srt.utils.tensor_bridge import use_mlx
|
||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||||
@@ -552,8 +552,8 @@ class Scheduler(
|
|||||||
self.grammar_manager = GrammarManager(self)
|
self.grammar_manager = GrammarManager(self)
|
||||||
|
|
||||||
self.request_receiver = SchedulerRequestReceiver(
|
self.request_receiver = SchedulerRequestReceiver(
|
||||||
recv_from_tokenizer=self.recv_from_tokenizer,
|
recv_from_tokenizer=self.ipc_channels.recv_from_tokenizer,
|
||||||
recv_from_rpc=self.recv_from_rpc,
|
recv_from_rpc=self.ipc_channels.recv_from_rpc,
|
||||||
recv_skipper=self.recv_skipper,
|
recv_skipper=self.recv_skipper,
|
||||||
input_blocker=self.input_blocker,
|
input_blocker=self.input_blocker,
|
||||||
mm_receiver=self.mm_receiver,
|
mm_receiver=self.mm_receiver,
|
||||||
@@ -629,7 +629,7 @@ class Scheduler(
|
|||||||
attn_dp_rank=self.ps.attn_dp_rank,
|
attn_dp_rank=self.ps.attn_dp_rank,
|
||||||
dp_rank=self.ps.dp_rank,
|
dp_rank=self.ps.dp_rank,
|
||||||
tree_cache=self.tree_cache,
|
tree_cache=self.tree_cache,
|
||||||
send_metrics_from_scheduler=self.send_metrics_from_scheduler,
|
send_metrics_from_scheduler=self.ipc_channels.send_metrics_from_scheduler,
|
||||||
max_running_requests=self.max_running_requests,
|
max_running_requests=self.max_running_requests,
|
||||||
max_total_num_tokens=self.max_total_num_tokens,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
get_stats=lambda: self.metrics_reporter.stats,
|
get_stats=lambda: self.metrics_reporter.stats,
|
||||||
@@ -658,7 +658,7 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.output_streamer = SchedulerOutputStreamer(
|
self.output_streamer = SchedulerOutputStreamer(
|
||||||
send_to_detokenizer=self.send_to_detokenizer,
|
send_to_detokenizer=self.ipc_channels.send_to_detokenizer,
|
||||||
tree_cache=self.tree_cache,
|
tree_cache=self.tree_cache,
|
||||||
ps=self.ps,
|
ps=self.ps,
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
@@ -726,50 +726,21 @@ class Scheduler(
|
|||||||
self.page_size = self.dllm_config.block_size
|
self.page_size = self.dllm_config.block_size
|
||||||
|
|
||||||
def init_ipc_channels(self, port_args: PortArgs):
|
def init_ipc_channels(self, port_args: PortArgs):
|
||||||
context = zmq.Context(2)
|
is_rank_zero = (
|
||||||
self.send_metrics_from_scheduler = None
|
|
||||||
|
|
||||||
if (
|
|
||||||
self.ps.pp_rank == 0
|
self.ps.pp_rank == 0
|
||||||
and self.ps.attn_tp_rank == 0
|
and self.ps.attn_tp_rank == 0
|
||||||
and self.ps.attn_cp_rank == 0
|
and self.ps.attn_cp_rank == 0
|
||||||
):
|
)
|
||||||
self.recv_from_tokenizer = get_zmq_socket(
|
self.ipc_channels = SchedulerIpcChannels.create(
|
||||||
context, zmq.PULL, port_args.scheduler_input_ipc_name, False
|
port_args=port_args,
|
||||||
)
|
is_rank_zero=is_rank_zero,
|
||||||
self.recv_from_rpc = get_zmq_socket(
|
skip_tokenizer_init=self.server_args.skip_tokenizer_init,
|
||||||
context, zmq.DEALER, port_args.rpc_ipc_name, False
|
metrics_enabled=self.server_args.enable_metrics
|
||||||
)
|
and (
|
||||||
|
self.ps.attn_tp_rank == 0
|
||||||
send_to_tokenizer = get_zmq_socket(
|
or self.server_args.enable_metrics_for_all_schedulers
|
||||||
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
),
|
||||||
)
|
)
|
||||||
if self.server_args.skip_tokenizer_init:
|
|
||||||
# Directly send to the TokenizerManager
|
|
||||||
send_to_detokenizer = get_zmq_socket(
|
|
||||||
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# Send to the DetokenizerManager
|
|
||||||
send_to_detokenizer = get_zmq_socket(
|
|
||||||
context, zmq.PUSH, port_args.detokenizer_ipc_name, False
|
|
||||||
)
|
|
||||||
|
|
||||||
self.send_to_tokenizer = SenderWrapper(send_to_tokenizer)
|
|
||||||
self.send_to_detokenizer = SenderWrapper(send_to_detokenizer)
|
|
||||||
else:
|
|
||||||
self.recv_from_tokenizer = None
|
|
||||||
self.recv_from_rpc = None
|
|
||||||
self.send_to_tokenizer = SenderWrapper(None)
|
|
||||||
self.send_to_detokenizer = SenderWrapper(None)
|
|
||||||
|
|
||||||
if self.server_args.enable_metrics and (
|
|
||||||
self.ps.attn_tp_rank == 0
|
|
||||||
or self.server_args.enable_metrics_for_all_schedulers
|
|
||||||
):
|
|
||||||
self.send_metrics_from_scheduler = get_zmq_socket(
|
|
||||||
context, zmq.PUSH, port_args.metrics_ipc_name, False
|
|
||||||
)
|
|
||||||
|
|
||||||
def init_idle_sleeper(self) -> None:
|
def init_idle_sleeper(self) -> None:
|
||||||
if (
|
if (
|
||||||
@@ -780,8 +751,8 @@ class Scheduler(
|
|||||||
):
|
):
|
||||||
self.idle_sleeper = IdleSleeper(
|
self.idle_sleeper = IdleSleeper(
|
||||||
sockets=[
|
sockets=[
|
||||||
self.recv_from_tokenizer,
|
self.ipc_channels.recv_from_tokenizer,
|
||||||
self.recv_from_rpc,
|
self.ipc_channels.recv_from_rpc,
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -919,7 +890,7 @@ class Scheduler(
|
|||||||
|
|
||||||
self.external_corpus_manager = ExternalCorpusManager(
|
self.external_corpus_manager = ExternalCorpusManager(
|
||||||
self.draft_worker,
|
self.draft_worker,
|
||||||
self.send_to_tokenizer.send_output,
|
self.ipc_channels.send_to_tokenizer.send_output,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.external_corpus_manager = None
|
self.external_corpus_manager = None
|
||||||
@@ -1660,10 +1631,10 @@ class Scheduler(
|
|||||||
output = self._request_dispatcher(recv_req)
|
output = self._request_dispatcher(recv_req)
|
||||||
if output is not None:
|
if output is not None:
|
||||||
if not isinstance(output, RpcReqOutput):
|
if not isinstance(output, RpcReqOutput):
|
||||||
self.send_to_tokenizer.send_output(output, recv_req)
|
self.ipc_channels.send_to_tokenizer.send_output(output, recv_req)
|
||||||
else:
|
else:
|
||||||
if self.recv_from_rpc is not None:
|
if self.ipc_channels.recv_from_rpc is not None:
|
||||||
self.recv_from_rpc.send_pyobj(output)
|
self.ipc_channels.recv_from_rpc.send_pyobj(output)
|
||||||
|
|
||||||
self._check_pending_flush()
|
self._check_pending_flush()
|
||||||
if self.external_corpus_manager is not None:
|
if self.external_corpus_manager is not None:
|
||||||
@@ -2093,7 +2064,7 @@ class Scheduler(
|
|||||||
rid=req.rid,
|
rid=req.rid,
|
||||||
)
|
)
|
||||||
req.time_stats.trace_ctx.abort(abort_info=abort_req.finished_reason)
|
req.time_stats.trace_ctx.abort(abort_info=abort_req.finished_reason)
|
||||||
self.send_to_tokenizer.send_output(abort_req, req)
|
self.ipc_channels.send_to_tokenizer.send_output(abort_req, req)
|
||||||
return False
|
return False
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -2132,7 +2103,7 @@ class Scheduler(
|
|||||||
req_to_abort = candidate_req
|
req_to_abort = candidate_req
|
||||||
message = "The request is aborted by a higher priority request."
|
message = "The request is aborted by a higher priority request."
|
||||||
|
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
AbortReq(
|
AbortReq(
|
||||||
finished_reason={
|
finished_reason={
|
||||||
"type": "abort",
|
"type": "abort",
|
||||||
@@ -2158,7 +2129,7 @@ class Scheduler(
|
|||||||
if self.enable_hicache_storage:
|
if self.enable_hicache_storage:
|
||||||
# Release prefetch events associated with the request
|
# Release prefetch events associated with the request
|
||||||
self.tree_cache.release_aborted_request(req.rid)
|
self.tree_cache.release_aborted_request(req.rid)
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
AbortReq(
|
AbortReq(
|
||||||
finished_reason={
|
finished_reason={
|
||||||
"type": "abort",
|
"type": "abort",
|
||||||
@@ -2750,7 +2721,7 @@ class Scheduler(
|
|||||||
self.new_token_ratio = new_token_ratio
|
self.new_token_ratio = new_token_ratio
|
||||||
for req in reqs_to_abort:
|
for req in reqs_to_abort:
|
||||||
abort_reason: FINISH_ABORT = req.to_finish
|
abort_reason: FINISH_ABORT = req.to_finish
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
AbortReq(
|
AbortReq(
|
||||||
finished_reason=abort_reason.to_json(),
|
finished_reason=abort_reason.to_json(),
|
||||||
rid=req.rid,
|
rid=req.rid,
|
||||||
@@ -2971,7 +2942,7 @@ class Scheduler(
|
|||||||
tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy()
|
tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy()
|
||||||
tp_active_ranks &= tp_active_ranks_cpu
|
tp_active_ranks &= tp_active_ranks_cpu
|
||||||
dp_active_ranks = tp_active_ranks.reshape(self.ps.dp_size, -1).prod(axis=1)
|
dp_active_ranks = tp_active_ranks.reshape(self.ps.dp_size, -1).prod(axis=1)
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
ActiveRanksOutput(status=dp_active_ranks.tolist())
|
ActiveRanksOutput(status=dp_active_ranks.tolist())
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -3039,7 +3010,7 @@ class Scheduler(
|
|||||||
# Return some signal for the health check.
|
# Return some signal for the health check.
|
||||||
# This is used to prevent the health check signal being blocked by long context prefill.
|
# This is used to prevent the health check signal being blocked by long context prefill.
|
||||||
# However, one minor issue is that this code path does not check the status of detokenizer manager.
|
# However, one minor issue is that this code path does not check the status of detokenizer manager.
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
HealthCheckOutput(
|
HealthCheckOutput(
|
||||||
http_worker_ipc=self.return_health_check_ipcs.popleft()
|
http_worker_ipc=self.return_health_check_ipcs.popleft()
|
||||||
)
|
)
|
||||||
@@ -3054,7 +3025,7 @@ class Scheduler(
|
|||||||
if self.is_fully_idle():
|
if self.is_fully_idle():
|
||||||
success = self.flush_cache()
|
success = self.flush_cache()
|
||||||
self._pending_flush = None
|
self._pending_flush = None
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
FlushCacheReqOutput(success=success), pending_req
|
FlushCacheReqOutput(success=success), pending_req
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
@@ -3064,7 +3035,7 @@ class Scheduler(
|
|||||||
"Deferred flush_cache timed out while waiting for idle state."
|
"Deferred flush_cache timed out while waiting for idle state."
|
||||||
)
|
)
|
||||||
self._pending_flush = None
|
self._pending_flush = None
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
FlushCacheReqOutput(
|
FlushCacheReqOutput(
|
||||||
success=False, message="Timed out waiting for idle state."
|
success=False, message="Timed out waiting for idle state."
|
||||||
),
|
),
|
||||||
@@ -3461,7 +3432,7 @@ class Scheduler(
|
|||||||
if self.enable_hicache_storage:
|
if self.enable_hicache_storage:
|
||||||
# to release prefetch events associated with the request
|
# to release prefetch events associated with the request
|
||||||
self.tree_cache.release_aborted_request(req.rid)
|
self.tree_cache.release_aborted_request(req.rid)
|
||||||
self.send_to_tokenizer.send_output(AbortReq(rid=req.rid), req)
|
self.ipc_channels.send_to_tokenizer.send_output(AbortReq(rid=req.rid), req)
|
||||||
# For disaggregation decode mode, the request in the waiting queue has KV cache allocated.
|
# For disaggregation decode mode, the request in the waiting queue has KV cache allocated.
|
||||||
if self.disaggregation_mode == DisaggregationMode.DECODE:
|
if self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||||
release_kv_cache(req, self.tree_cache)
|
release_kv_cache(req, self.tree_cache)
|
||||||
@@ -3525,7 +3496,7 @@ class Scheduler(
|
|||||||
if recv_req.abort_all or decode_req.rid.startswith(recv_req.rid):
|
if recv_req.abort_all or decode_req.rid.startswith(recv_req.rid):
|
||||||
assert hasattr(decode_req, "kv_cache_cpu")
|
assert hasattr(decode_req, "kv_cache_cpu")
|
||||||
del decode_req.kv_cache_cpu
|
del decode_req.kv_cache_cpu
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
AbortReq(rid=decode_req.rid), decode_req
|
AbortReq(rid=decode_req.rid), decode_req
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -3687,7 +3658,7 @@ class Scheduler(
|
|||||||
def handle_freeze_gc(self, recv_req: FreezeGCReq):
|
def handle_freeze_gc(self, recv_req: FreezeGCReq):
|
||||||
"""Handle freeze_gc request: freeze scheduler's GC and forward to detokenizer."""
|
"""Handle freeze_gc request: freeze scheduler's GC and forward to detokenizer."""
|
||||||
freeze_gc("Scheduler")
|
freeze_gc("Scheduler")
|
||||||
self.send_to_detokenizer.send_output(recv_req, recv_req)
|
self.ipc_channels.send_to_detokenizer.send_output(recv_req, recv_req)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def handle_dumper_control(self, recv_req: DumperControlReqInput):
|
def handle_dumper_control(self, recv_req: DumperControlReqInput):
|
||||||
@@ -3702,12 +3673,12 @@ class Scheduler(
|
|||||||
response = dumper._http_manager.handle_request(
|
response = dumper._http_manager.handle_request(
|
||||||
method=recv_req.method, body=recv_req.body
|
method=recv_req.method, body=recv_req.body
|
||||||
)
|
)
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
DumperControlReqOutput(success=True, response=response), recv_req
|
DumperControlReqOutput(success=True, response=response), recv_req
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"[Scheduler] handle_dumper_control error: {e}", flush=True)
|
print(f"[Scheduler] handle_dumper_control error: {e}", flush=True)
|
||||||
self.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
DumperControlReqOutput(success=False, response=[], error=str(e)),
|
DumperControlReqOutput(success=False, response=[], error=str(e)),
|
||||||
recv_req,
|
recv_req,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,73 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import zmq
|
||||||
|
|
||||||
|
from sglang.srt.managers.scheduler_components.output_sender import SenderWrapper
|
||||||
|
from sglang.srt.server_args import PortArgs
|
||||||
|
from sglang.srt.utils.network import get_zmq_socket
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||||
|
class SchedulerIpcChannels:
|
||||||
|
recv_from_tokenizer: Optional[zmq.Socket]
|
||||||
|
recv_from_rpc: Optional[zmq.Socket]
|
||||||
|
send_to_tokenizer: SenderWrapper
|
||||||
|
send_to_detokenizer: SenderWrapper
|
||||||
|
send_metrics_from_scheduler: Optional[zmq.Socket]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
*,
|
||||||
|
port_args: PortArgs,
|
||||||
|
is_rank_zero: bool,
|
||||||
|
skip_tokenizer_init: bool,
|
||||||
|
metrics_enabled: bool,
|
||||||
|
) -> "SchedulerIpcChannels":
|
||||||
|
context = zmq.Context(2)
|
||||||
|
|
||||||
|
if is_rank_zero:
|
||||||
|
recv_from_tokenizer = get_zmq_socket(
|
||||||
|
context, zmq.PULL, port_args.scheduler_input_ipc_name, False
|
||||||
|
)
|
||||||
|
recv_from_rpc = get_zmq_socket(
|
||||||
|
context, zmq.DEALER, port_args.rpc_ipc_name, False
|
||||||
|
)
|
||||||
|
|
||||||
|
send_to_tokenizer_raw = get_zmq_socket(
|
||||||
|
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
||||||
|
)
|
||||||
|
if skip_tokenizer_init:
|
||||||
|
# Directly send to the TokenizerManager
|
||||||
|
send_to_detokenizer_raw = get_zmq_socket(
|
||||||
|
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Send to the DetokenizerManager
|
||||||
|
send_to_detokenizer_raw = get_zmq_socket(
|
||||||
|
context, zmq.PUSH, port_args.detokenizer_ipc_name, False
|
||||||
|
)
|
||||||
|
|
||||||
|
send_to_tokenizer = SenderWrapper(send_to_tokenizer_raw)
|
||||||
|
send_to_detokenizer = SenderWrapper(send_to_detokenizer_raw)
|
||||||
|
else:
|
||||||
|
recv_from_tokenizer = None
|
||||||
|
recv_from_rpc = None
|
||||||
|
send_to_tokenizer = SenderWrapper(None)
|
||||||
|
send_to_detokenizer = SenderWrapper(None)
|
||||||
|
|
||||||
|
if metrics_enabled:
|
||||||
|
send_metrics_from_scheduler = get_zmq_socket(
|
||||||
|
context, zmq.PUSH, port_args.metrics_ipc_name, False
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
send_metrics_from_scheduler = None
|
||||||
|
|
||||||
|
return cls(
|
||||||
|
recv_from_tokenizer=recv_from_tokenizer,
|
||||||
|
recv_from_rpc=recv_from_rpc,
|
||||||
|
send_to_tokenizer=send_to_tokenizer,
|
||||||
|
send_to_detokenizer=send_to_detokenizer,
|
||||||
|
send_metrics_from_scheduler=send_metrics_from_scheduler,
|
||||||
|
)
|
||||||
@@ -30,7 +30,7 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase):
|
|||||||
scheduler.model_config = SimpleNamespace(num_key_value_heads=8)
|
scheduler.model_config = SimpleNamespace(num_key_value_heads=8)
|
||||||
scheduler.disagg_prefill_bootstrap_queue = MagicMock()
|
scheduler.disagg_prefill_bootstrap_queue = MagicMock()
|
||||||
scheduler.disagg_decode_prealloc_queue = MagicMock()
|
scheduler.disagg_decode_prealloc_queue = MagicMock()
|
||||||
scheduler.send_to_tokenizer = MagicMock()
|
scheduler.ipc_channels = MagicMock()
|
||||||
return scheduler
|
return scheduler
|
||||||
|
|
||||||
def _new_req(self, priority=None):
|
def _new_req(self, priority=None):
|
||||||
@@ -72,7 +72,7 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase):
|
|||||||
scheduler._add_request_to_queue(req)
|
scheduler._add_request_to_queue(req)
|
||||||
|
|
||||||
scheduler.disagg_decode_prealloc_queue.add.assert_not_called()
|
scheduler.disagg_decode_prealloc_queue.add.assert_not_called()
|
||||||
scheduler.send_to_tokenizer.send_output.assert_called_once()
|
scheduler.ipc_channels.send_to_tokenizer.send_output.assert_called_once()
|
||||||
req.time_stats.trace_ctx.abort.assert_called_once()
|
req.time_stats.trace_ctx.abort.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
|||||||
def _new_scheduler(self) -> Scheduler:
|
def _new_scheduler(self) -> Scheduler:
|
||||||
scheduler = Scheduler.__new__(Scheduler)
|
scheduler = Scheduler.__new__(Scheduler)
|
||||||
scheduler._pending_flush = None
|
scheduler._pending_flush = None
|
||||||
scheduler.send_to_tokenizer = MagicMock()
|
scheduler.ipc_channels = MagicMock()
|
||||||
scheduler.flush_cache = MagicMock(return_value=True)
|
scheduler.flush_cache = MagicMock(return_value=True)
|
||||||
scheduler.is_fully_idle = MagicMock(return_value=False)
|
scheduler.is_fully_idle = MagicMock(return_value=False)
|
||||||
return scheduler
|
return scheduler
|
||||||
@@ -82,7 +82,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertIsNone(scheduler._pending_flush)
|
self.assertIsNone(scheduler._pending_flush)
|
||||||
scheduler.flush_cache.assert_called_once()
|
scheduler.flush_cache.assert_called_once()
|
||||||
out = scheduler.send_to_tokenizer.send_output.call_args.args[0]
|
out = scheduler.ipc_channels.send_to_tokenizer.send_output.call_args.args[0]
|
||||||
self.assertTrue(out.success)
|
self.assertTrue(out.success)
|
||||||
|
|
||||||
def test_pending_flush_expires_on_timeout(self):
|
def test_pending_flush_expires_on_timeout(self):
|
||||||
@@ -95,7 +95,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertIsNone(scheduler._pending_flush)
|
self.assertIsNone(scheduler._pending_flush)
|
||||||
scheduler.flush_cache.assert_not_called()
|
scheduler.flush_cache.assert_not_called()
|
||||||
out = scheduler.send_to_tokenizer.send_output.call_args.args[0]
|
out = scheduler.ipc_channels.send_to_tokenizer.send_output.call_args.args[0]
|
||||||
self.assertFalse(out.success)
|
self.assertFalse(out.success)
|
||||||
|
|
||||||
def test_pending_flush_survives_before_deadline(self):
|
def test_pending_flush_survives_before_deadline(self):
|
||||||
@@ -107,7 +107,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
|||||||
Scheduler._check_pending_flush(scheduler)
|
Scheduler._check_pending_flush(scheduler)
|
||||||
|
|
||||||
self.assertIsNotNone(scheduler._pending_flush)
|
self.assertIsNotNone(scheduler._pending_flush)
|
||||||
scheduler.send_to_tokenizer.send_output.assert_not_called()
|
scheduler.ipc_channels.send_to_tokenizer.send_output.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user