From db7a751d4869b516ab3510919a597bf1815d534f Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 16 Apr 2026 15:37:19 -0700 Subject: [PATCH] refactor: extract FanOutCommunicator and use declarative spec table (#22967) --- .../hicache_storage_runtime_attach_detach.md | 4 +- python/sglang/srt/managers/communicator.py | 93 +++++ .../srt/managers/multi_tokenizer_mixin.py | 4 +- ...or_mixin.py => tokenizer_control_mixin.py} | 349 +++--------------- .../sglang/srt/managers/tokenizer_manager.py | 4 +- .../unit/utils/test_profile_merger.py | 6 +- 6 files changed, 161 insertions(+), 299 deletions(-) create mode 100644 python/sglang/srt/managers/communicator.py rename python/sglang/srt/managers/{tokenizer_communicator_mixin.py => tokenizer_control_mixin.py} (74%) diff --git a/docs/advanced_features/hicache_storage_runtime_attach_detach.md b/docs/advanced_features/hicache_storage_runtime_attach_detach.md index deaef8c12..555d799c2 100644 --- a/docs/advanced_features/hicache_storage_runtime_attach_detach.md +++ b/docs/advanced_features/hicache_storage_runtime_attach_detach.md @@ -19,8 +19,8 @@ The control path is: 1. **HTTP Server** (`python/sglang/srt/entrypoints/http_server.py`) - Exposes `PUT /hicache/storage-backend`, `DELETE /hicache/storage-backend`, `GET /hicache/storage-backend` -2. **TokenizerManager** (`python/sglang/srt/managers/tokenizer_communicator_mixin.py`) - - Sends the request to the Scheduler via `_Communicator` +2. **TokenizerManager** (`python/sglang/srt/managers/tokenizer_control_mixin.py`) + - Sends the request to the Scheduler via `FanOutCommunicator` 3. **Scheduler** (`python/sglang/srt/managers/scheduler.py`) - Performs a **strict idle check** - Calls `tree_cache.attach_storage_backend(...)` / `detach_storage_backend(...)` diff --git a/python/sglang/srt/managers/communicator.py b/python/sglang/srt/managers/communicator.py new file mode 100644 index 000000000..3080f6a75 --- /dev/null +++ b/python/sglang/srt/managers/communicator.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +import asyncio +import copy +from collections import deque +from typing import Deque, Generic, List, Optional, TypeVar + +import zmq + +T = TypeVar("T") + + +class FanOutCommunicator(Generic[T]): + """Fan-out request + collect response primitive over zmq. + + One send is fanned out to `fan_out` recipients; the caller awaits until + all `fan_out` responses are collected. Supports two modes: + - "queueing": requests are serialized; concurrent callers wait in a FIFO queue. + - "watching": concurrent callers share a single in-flight request and all + receive the same result when it completes. + + Only one request is in-flight at any time in either mode. + """ + + def __init__(self, sender: zmq.Socket, fan_out: int, mode="queueing"): + self._sender = sender + self._fan_out = fan_out + self._mode = mode + self._result_event: Optional[asyncio.Event] = None + self._result_values: Optional[List[T]] = None + self._ready_queue: Deque[asyncio.Event] = deque() + + assert mode in ["queueing", "watching"] + + async def queueing_call(self, obj: T): + ready_event = asyncio.Event() + if self._result_event is not None or len(self._ready_queue) > 0: + self._ready_queue.append(ready_event) + await ready_event.wait() + assert self._result_event is None + assert self._result_values is None + + if obj is not None: + self._sender.send_pyobj(obj) + + self._result_event = asyncio.Event() + self._result_values = [] + await self._result_event.wait() + result_values = self._result_values + self._result_event = self._result_values = None + + if len(self._ready_queue) > 0: + self._ready_queue.popleft().set() + + return result_values + + async def watching_call(self, obj): + if self._result_event is None: + assert self._result_values is None + self._result_values = [] + self._result_event = asyncio.Event() + + if obj is not None: + self._sender.send_pyobj(obj) + + # Capture local refs before await -- after event fires, the first + # awakened coroutine clears shared state; later awaiters use local refs. + values = self._result_values + event = self._result_event + await event.wait() + + result_values = copy.deepcopy(values) + if self._result_event is event: + self._result_event = self._result_values = None + return result_values + + async def __call__(self, obj): + if self._mode == "queueing": + return await self.queueing_call(obj) + else: + return await self.watching_call(obj) + + def handle_recv(self, recv_obj: T): + self._result_values.append(recv_obj) + if len(self._result_values) == self._fan_out: + self._result_event.set() + + @staticmethod + def merge_results(results): + all_success = all([r.success for r in results]) + all_message = [r.message for r in results] + all_message = " | ".join(all_message) + return all_success, all_message diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 8da3d3b0d..9757e4b58 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -35,6 +35,7 @@ import zmq import zmq.asyncio from sglang.srt.disaggregation.utils import DisaggregationMode, TransferBackend +from sglang.srt.managers.communicator import FanOutCommunicator from sglang.srt.managers.disagg_service import start_disagg_service from sglang.srt.managers.io_struct import ( BaseBatchReq, @@ -43,7 +44,6 @@ from sglang.srt.managers.io_struct import ( BatchStrOutput, BatchTokenIDOutput, ) -from sglang.srt.managers.tokenizer_communicator_mixin import _Communicator from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import kill_process_tree @@ -400,7 +400,7 @@ class TokenizerWorker(TokenizerManager): self.server_args.disaggregation_transfer_backend ) # Communicator - self.register_multi_tokenizer_communicator = _Communicator( + self.register_multi_tokenizer_communicator = FanOutCommunicator( self.send_to_scheduler, 2 ) diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py similarity index 74% rename from python/sglang/srt/managers/tokenizer_communicator_mixin.py rename to python/sglang/srt/managers/tokenizer_control_mixin.py index 9aab35179..b18c3cd34 100644 --- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -1,26 +1,21 @@ from __future__ import annotations import asyncio -import copy import logging import time import uuid -from collections import deque from typing import ( TYPE_CHECKING, Any, - Deque, Dict, - Generic, List, Optional, Tuple, - TypeVar, ) import fastapi -import zmq +from sglang.srt.managers.communicator import FanOutCommunicator from sglang.srt.managers.io_struct import ( AddExternalCorpusReqInput, AddExternalCorpusReqOutput, @@ -93,286 +88,60 @@ from sglang.utils import TypeBasedDispatcher if TYPE_CHECKING: from sglang.srt.managers.tokenizer_manager import TokenizerManager -T = TypeVar("T") - logger = logging.getLogger(__name__) - -class _Communicator(Generic[T]): - """Note: The communicator now only run up to 1 in-flight request at any time.""" - - def __init__(self, sender: zmq.Socket, fan_out: int, mode="queueing"): - self._sender = sender - self._fan_out = fan_out - self._mode = mode - self._result_event: Optional[asyncio.Event] = None - self._result_values: Optional[List[T]] = None - self._ready_queue: Deque[asyncio.Future] = deque() - - assert mode in ["queueing", "watching"] - - async def queueing_call(self, obj: T): - ready_event = asyncio.Event() - if self._result_event is not None or len(self._ready_queue) > 0: - self._ready_queue.append(ready_event) - await ready_event.wait() - assert self._result_event is None - assert self._result_values is None - - if obj: - self._sender.send_pyobj(obj) - - self._result_event = asyncio.Event() - self._result_values = [] - await self._result_event.wait() - result_values = self._result_values - self._result_event = self._result_values = None - - if len(self._ready_queue) > 0: - self._ready_queue.popleft().set() - - return result_values - - async def watching_call(self, obj): - if self._result_event is None: - assert self._result_values is None - self._result_values = [] - self._result_event = asyncio.Event() - - if obj: - self._sender.send_pyobj(obj) - - # NOTE: Capture list ref before await so later awaiters survive clearing. - values = self._result_values - event = self._result_event - await event.wait() - - result_values = copy.deepcopy(values) - if self._result_event is event: - self._result_event = self._result_values = None - return result_values - - async def __call__(self, obj): - if self._mode == "queueing": - return await self.queueing_call(obj) - else: - return await self.watching_call(obj) - - def handle_recv(self, recv_obj: T): - self._result_values.append(recv_obj) - if len(self._result_values) == self._fan_out: - self._result_event.set() - - @staticmethod - def merge_results(results): - all_success = all([r.success for r in results]) - all_message = [r.message for r in results] - all_message = " | ".join(all_message) - return all_success, all_message +# Declarative spec: (attr_name_prefix, response_type[, mode]) +# Each entry creates self.{prefix}_communicator and registers +# response_type -> communicator.handle_recv in the dispatch table. +_COMMUNICATOR_SPECS = [ + ("init_weights_update_group", InitWeightsUpdateGroupReqOutput), + ("destroy_weights_update_group", DestroyWeightsUpdateGroupReqOutput), + ("update_weights_from_distributed", UpdateWeightsFromDistributedReqOutput), + ( + "init_weights_send_group_for_remote_instance", + InitWeightsSendGroupForRemoteInstanceReqOutput, + ), + ("send_weights_to_remote_instance", SendWeightsToRemoteInstanceReqOutput), + ("update_weights_from_tensor", UpdateWeightsFromTensorReqOutput), + ("update_weights_from_ipc", UpdateWeightsFromIPCReqOutput), + ("get_weights_by_name", GetWeightsByNameReqOutput), + ("release_memory_occupation", ReleaseMemoryOccupationReqOutput), + ("resume_memory_occupation", ResumeMemoryOccupationReqOutput), + ("check_weights", CheckWeightsReqOutput), + ("slow_down", SlowDownReqOutput), + ("flush_cache", FlushCacheReqOutput), + ("add_external_corpus", AddExternalCorpusReqOutput), + ("remove_external_corpus", RemoveExternalCorpusReqOutput), + ("list_external_corpora", ListExternalCorporaReqOutput), + ("clear_hicache_storage", ClearHiCacheReqOutput), + ("attach_hicache_storage", AttachHiCacheStorageReqOutput), + ("detach_hicache_storage", DetachHiCacheStorageReqOutput), + ("profile", ProfileReqOutput), + ("get_internal_state", GetInternalStateReqOutput), + ("set_internal_state", SetInternalStateReqOutput), + ("expert_distribution", ExpertDistributionReqOutput), + ("update_lora_adapter", LoRAUpdateOutput), + ("get_load", GetLoadReqOutput, "watching"), + ("get_loads", GetLoadsReqOutput, "watching"), + ("dumper_control", DumperControlReqOutput), +] -class TokenizerCommunicatorMixin: - """Mixin class for TokenizerManager to handle communication with the scheduler.""" +class TokenizerControlMixin: + """Mixin for TokenizerManager's control-plane operations (weights, cache, lora, + profile, internal state, etc.) -- everything that talks to the scheduler via + FanOutCommunicator, as opposed to data-plane inference requests multiplexed by rid. + """ def init_communicators(self: TokenizerManager, server_args: ServerArgs): - # Communicators - self.init_weights_update_group_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.destroy_weights_update_group_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.update_weights_from_distributed_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.init_weights_send_group_for_remote_instance_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.send_weights_to_remote_instance_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.update_weights_from_tensor_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.update_weights_from_ipc_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.get_weights_by_name_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.release_memory_occupation_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.resume_memory_occupation_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.check_weights_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.slow_down_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.flush_cache_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.add_external_corpus_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.remove_external_corpus_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.list_external_corpora_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.clear_hicache_storage_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.attach_hicache_storage_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.detach_hicache_storage_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.profile_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.get_internal_state_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.set_internal_state_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.expert_distribution_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.update_lora_adapter_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - self.get_load_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size, mode="watching" - ) - self.get_loads_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size, mode="watching" - ) - self.dumper_control_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) - - self._result_dispatcher += self._get_communicator_dispatcher() - - def _get_communicator_dispatcher(self: TokenizerManager): - return TypeBasedDispatcher( - [ - ( - InitWeightsUpdateGroupReqOutput, - self.init_weights_update_group_communicator.handle_recv, - ), - ( - DestroyWeightsUpdateGroupReqOutput, - self.destroy_weights_update_group_communicator.handle_recv, - ), - ( - UpdateWeightsFromDistributedReqOutput, - self.update_weights_from_distributed_communicator.handle_recv, - ), - ( - InitWeightsSendGroupForRemoteInstanceReqOutput, - self.init_weights_send_group_for_remote_instance_communicator.handle_recv, - ), - ( - SendWeightsToRemoteInstanceReqOutput, - self.send_weights_to_remote_instance_communicator.handle_recv, - ), - ( - UpdateWeightsFromTensorReqOutput, - self.update_weights_from_tensor_communicator.handle_recv, - ), - ( - UpdateWeightsFromIPCReqOutput, - self.update_weights_from_ipc_communicator.handle_recv, - ), - ( - GetWeightsByNameReqOutput, - self.get_weights_by_name_communicator.handle_recv, - ), - ( - ReleaseMemoryOccupationReqOutput, - self.release_memory_occupation_communicator.handle_recv, - ), - ( - ResumeMemoryOccupationReqOutput, - self.resume_memory_occupation_communicator.handle_recv, - ), - ( - CheckWeightsReqOutput, - self.check_weights_communicator.handle_recv, - ), - ( - SlowDownReqOutput, - self.slow_down_communicator.handle_recv, - ), - ( - ClearHiCacheReqOutput, - self.clear_hicache_storage_communicator.handle_recv, - ), - ( - AttachHiCacheStorageReqOutput, - self.attach_hicache_storage_communicator.handle_recv, - ), - ( - DetachHiCacheStorageReqOutput, - self.detach_hicache_storage_communicator.handle_recv, - ), - ( - FlushCacheReqOutput, - self.flush_cache_communicator.handle_recv, - ), - ( - AddExternalCorpusReqOutput, - self.add_external_corpus_communicator.handle_recv, - ), - ( - RemoveExternalCorpusReqOutput, - self.remove_external_corpus_communicator.handle_recv, - ), - ( - ListExternalCorporaReqOutput, - self.list_external_corpora_communicator.handle_recv, - ), - ( - ProfileReqOutput, - self.profile_communicator.handle_recv, - ), - ( - GetInternalStateReqOutput, - self.get_internal_state_communicator.handle_recv, - ), - ( - SetInternalStateReqOutput, - self.set_internal_state_communicator.handle_recv, - ), - ( - ExpertDistributionReqOutput, - self.expert_distribution_communicator.handle_recv, - ), - ( - LoRAUpdateOutput, - self.update_lora_adapter_communicator.handle_recv, - ), - ( - GetLoadReqOutput, - self.get_load_communicator.handle_recv, - ), - ( - GetLoadsReqOutput, - self.get_loads_communicator.handle_recv, - ), - ( - DumperControlReqOutput, - self.dumper_control_communicator.handle_recv, - ), - ] - ) + dispatch_pairs = [] + for spec in _COMMUNICATOR_SPECS: + name, resp_type = spec[0], spec[1] + mode = spec[2] if len(spec) > 2 else "queueing" + comm = FanOutCommunicator(self.send_to_scheduler, server_args.dp_size, mode) + setattr(self, f"{name}_communicator", comm) + dispatch_pairs.append((resp_type, comm.handle_recv)) + self._result_dispatcher += TypeBasedDispatcher(dispatch_pairs) async def add_external_corpus( self: TokenizerManager, obj: AddExternalCorpusReqInput @@ -438,7 +207,7 @@ class TokenizerCommunicatorMixin: obj.file_path = None obj.documents = None results = await self.add_external_corpus_communicator(obj) - all_success, all_message = _Communicator.merge_results(results) + all_success, all_message = FanOutCommunicator.merge_results(results) if truncated and all_success: all_message += f" (truncated: exceeded {max_tokens} token limit)" return AddExternalCorpusReqOutput( @@ -462,7 +231,7 @@ class TokenizerCommunicatorMixin: results = await self.remove_external_corpus_communicator( RemoveExternalCorpusReqInput(corpus_id=corpus_id) ) - all_success, all_message = _Communicator.merge_results(results) + all_success, all_message = FanOutCommunicator.merge_results(results) return RemoveExternalCorpusReqOutput(success=all_success, message=all_message) async def list_external_corpora( @@ -477,7 +246,7 @@ class TokenizerCommunicatorMixin: results = await self.list_external_corpora_communicator( ListExternalCorporaReqInput() ) - all_success, all_message = _Communicator.merge_results(results) + all_success, all_message = FanOutCommunicator.merge_results(results) # Merge corpus token counts from all DP ranks (each rank loads the same set). corpus_token_counts = results[0].corpus_token_counts if all_success else {} return ListExternalCorporaReqOutput( @@ -520,7 +289,7 @@ class TokenizerCommunicatorMixin: ) ) - all_success, all_message = _Communicator.merge_results(results) + all_success, all_message = FanOutCommunicator.merge_results(results) out = AttachHiCacheStorageReqOutput(success=all_success, message=all_message) # TODO: partial rollback if failed if all_success: @@ -547,7 +316,7 @@ class TokenizerCommunicatorMixin: DetachHiCacheStorageReqInput() ) - all_success, all_message = _Communicator.merge_results(results) + all_success, all_message = FanOutCommunicator.merge_results(results) out = DetachHiCacheStorageReqOutput(success=all_success, message=all_message) # TODO: partial rollback if failed if all_success: @@ -628,7 +397,7 @@ class TokenizerCommunicatorMixin: ), "dp_size must be 1 or dp attention must be enabled for update weights from distributed" results = await self.init_weights_update_group_communicator(obj) - return _Communicator.merge_results(results) + return FanOutCommunicator.merge_results(results) async def destroy_weights_update_group( self: TokenizerManager, @@ -641,7 +410,7 @@ class TokenizerCommunicatorMixin: ), "dp_size must be 1 or dp attention must be enabled for destroy parameter update group" results = await self.destroy_weights_update_group_communicator(obj) - return _Communicator.merge_results(results) + return FanOutCommunicator.merge_results(results) async def update_weights_from_distributed( self: TokenizerManager, @@ -666,7 +435,7 @@ class TokenizerCommunicatorMixin: async with self.model_update_lock.writer_lock: results = await self.update_weights_from_distributed_communicator(obj) - success, message = _Communicator.merge_results(results) + success, message = FanOutCommunicator.merge_results(results) if success and obj.weight_version is not None: self._update_weight_version_if_provided(obj.weight_version) message += f" Weight version updated to {obj.weight_version}." @@ -723,7 +492,7 @@ class TokenizerCommunicatorMixin: async with self.model_update_lock.writer_lock: results = await self.update_weights_from_tensor_communicator(obj) - success, message = _Communicator.merge_results(results) + success, message = FanOutCommunicator.merge_results(results) if success and obj.weight_version is not None: self._update_weight_version_if_provided(obj.weight_version) message += f" Weight version updated to {obj.weight_version}." @@ -1001,7 +770,7 @@ class TokenizerCommunicatorMixin: ) -> CheckWeightsReqOutput: self.auto_create_handle_loop() results = await self.check_weights_communicator(obj) - return _Communicator.merge_results(results) + return FanOutCommunicator.merge_results(results) async def slow_down( self: TokenizerManager, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 3968411ee..8a49717c5 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -76,7 +76,7 @@ from sglang.srt.managers.multimodal_processor import get_mm_processor, import_pr from sglang.srt.managers.schedule_batch import MultimodalDataItem from sglang.srt.managers.scheduler import is_health_check_generate_req from sglang.srt.managers.scheduler_input_blocker import input_blocker_guard_region -from sglang.srt.managers.tokenizer_communicator_mixin import TokenizerCommunicatorMixin +from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin from sglang.srt.managers.tokenizer_manager_score_mixin import ( TokenizerManagerScoreMixin, ) @@ -212,7 +212,7 @@ class InputFormat(Enum): CROSS_ENCODER_PAIRS = 3 # Cross-encoder pairs like [["query", "document"]] -class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin): +class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): """TokenizerManager is a process that tokenizes the text.""" def __init__( diff --git a/test/registered/unit/utils/test_profile_merger.py b/test/registered/unit/utils/test_profile_merger.py index 5c00e793e..e410af33e 100644 --- a/test/registered/unit/utils/test_profile_merger.py +++ b/test/registered/unit/utils/test_profile_merger.py @@ -221,11 +221,11 @@ class TestProfileMergerIntegration(unittest.TestCase): import inspect # Test TokenizerManager - from sglang.srt.managers.tokenizer_communicator_mixin import ( - TokenizerCommunicatorMixin, + from sglang.srt.managers.tokenizer_control_mixin import ( + TokenizerControlMixin, ) - sig = inspect.signature(TokenizerCommunicatorMixin.start_profile) + sig = inspect.signature(TokenizerControlMixin.start_profile) self.assertIn("merge_profiles", sig.parameters) # Test SchedulerProfilerMixin