From f25b067756459cb73e22c4a77588f828a329e2b3 Mon Sep 17 00:00:00 2001 From: zhangxiaolei Date: Tue, 21 Jul 2026 14:58:47 +0800 Subject: [PATCH 1/4] Move PD hidden capture state out of schedule batch --- python/sglang/srt/disaggregation/prefill.py | 23 +++---------------- python/sglang/srt/disaggregation/utils.py | 9 ++++++++ python/sglang/srt/managers/schedule_batch.py | 15 +----------- .../sglang/srt/managers/scheduler_pp_mixin.py | 10 ++++---- .../srt/model_executor/forward_batch_info.py | 6 +++-- 5 files changed, 23 insertions(+), 40 deletions(-) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index d921663b784f..8a1358f25f95 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -42,6 +42,7 @@ TransferBackend, get_dsv4_c128_state_indices, get_kv_class, + get_pd_hidden_capture_layer_ids, is_aborted, is_dsv4_c128_online_enabled, is_mla_backend, @@ -57,7 +58,6 @@ Req, ScheduleBatch, ) -from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode from sglang.srt.mem_cache.common import ( kv_to_page_indices, kv_to_page_num, @@ -831,30 +831,12 @@ def get_next_disagg_prefill_batch_to_run( batch = prefill_plan.batch_to_run running_batch = prefill_plan.running_batch batch = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch(batch) - self._prepare_pd_hidden_capture_for_batch(batch) if batch: set_schedule_time_batch(batch) return NextBatchPlan(batch_to_run=batch, running_batch=running_batch) - def _prepare_pd_hidden_capture_for_batch( - self: Scheduler, batch: Optional[ScheduleBatch] - ) -> None: - dspark_capture_layers = None - if batch: - for req in batch.reqs: - dspark_capture_layers = getattr( - req, "pd_hidden_capture_layer_ids", None - ) - if dspark_capture_layers: - break - if dspark_capture_layers: - batch.pd_hidden_capture_layer_ids = [ - int(x) for x in dspark_capture_layers - ] - batch.capture_hidden_mode = CaptureHiddenMode.FULL - @torch.no_grad() def event_loop_normal_disagg_prefill(self: Scheduler) -> None: """A normal scheduler loop for prefill worker in disaggregation mode.""" @@ -1054,7 +1036,8 @@ def _write_pd_hidden_rows_for_batch( ] message = ( "PD hidden capture was required but forward output has no " - f"hidden states: batch_capture_layers={batch.pd_hidden_capture_layer_ids}, " + "hidden states: batch_capture_layers=" + f"{get_pd_hidden_capture_layer_ids(batch.reqs)}, " f"reqs={reqs}" ) for req in needs_pd_hidden_reqs: diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 379a7a0e7032..8f039aef184b 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -80,6 +80,15 @@ def get_dsv4_c128_state_indices( return np.array([page], dtype=np.int32) +def get_pd_hidden_capture_layer_ids(reqs: List["Req"]) -> Optional[List[int]]: + """Return the per-batch PD hidden capture layers requested by any req.""" + for req in reqs: + layer_ids = getattr(req, "pd_hidden_capture_layer_ids", None) + if layer_ids: + return [int(x) for x in layer_ids] + return None + + class DisaggregationMode(Enum): NULL = "null" PREFILL = "prefill" diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 9576bedd4d8f..40d90adfe5ac 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1983,12 +1983,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # spec_info: Optional[SpecInput] = None spec_info: Optional[SpecInput] = None - # === One-shot per-forward overrides; init_new consumes and resets === - seq_lens_cpu_cache: torch.Tensor = None - capture_hidden_mode: Optional[CaptureHiddenMode] = None - return_hidden_states_before_norm: bool = False - pd_hidden_capture_layer_ids: Optional[List[int]] = None - @classmethod def init_new( cls, @@ -3032,9 +3026,7 @@ def copy(self): can_run_dp_breakable_cuda_graph=self.can_run_dp_breakable_cuda_graph, is_extend_in_batch=self.is_extend_in_batch, is_prefill_only=self.is_prefill_only, - seq_lens_cpu=( - self.seq_lens_cpu.clone() if self.seq_lens_cpu is not None else None - ), + seq_lens_cpu=self.seq_lens_cpu, enable_overlap=self.enable_overlap, mamba_track_indices=self.mamba_track_indices, mamba_track_mask=self.mamba_track_mask, @@ -3043,11 +3035,6 @@ def copy(self): prefill_stats=self.prefill_stats, fpm_start_time=self.fpm_start_time, forward_iter=self.forward_iter, - pd_hidden_capture_layer_ids=( - self.pd_hidden_capture_layer_ids[:] - if self.pd_hidden_capture_layer_ids is not None - else None - ), ) def maybe_evict_swa(self): diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 4fd1d6c29ebf..f4dab0962e8e 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -14,7 +14,10 @@ from tqdm import tqdm from sglang.srt.disaggregation.base.conn import KVPoll -from sglang.srt.disaggregation.utils import poll_and_all_reduce_attn_cp_tp_group +from sglang.srt.disaggregation.utils import ( + get_pd_hidden_capture_layer_ids, + poll_and_all_reduce_attn_cp_tp_group, +) from sglang.srt.distributed.parallel_state import P2PWork from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import ( @@ -259,7 +262,6 @@ def event_loop_pp_disagg_prefill(self: Scheduler): batch = prefill_plan.batch_to_run self.running_batch = prefill_plan.running_batch batch = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch(batch) - self._prepare_pd_hidden_capture_for_batch(batch) self.mbs[mb_id] = batch self.running_mbs[mb_id] = self.running_batch @@ -1016,7 +1018,7 @@ def _pp_prepare_tensor_dict( **logprob_dict, } if ( - batch.pd_hidden_capture_layer_ids + get_pd_hidden_capture_layer_ids(batch.reqs) and not self._pp_should_owner_direct_pd_hidden(batch) and result.logits_output is not None and result.logits_output.hidden_states is not None @@ -1029,7 +1031,7 @@ def _pp_should_owner_direct_pd_hidden( ) -> bool: if not hasattr(self, "disagg_prefill_bootstrap_queue"): return False - if not batch or not batch.pd_hidden_capture_layer_ids: + if not batch or not get_pd_hidden_capture_layer_ids(batch.reqs): return False capture_reqs = [ req diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 981b39ff55f7..c39a94b3f0e8 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -42,6 +42,7 @@ from sglang.srt.kv_canary.req_to_expected_token_ids_manager import ( compute_req_all_ids_info, ) +from sglang.srt.disaggregation.utils import get_pd_hidden_capture_layer_ids from sglang.srt.layers.dp_attention import ( DpPaddingMode, set_dp_buffer_len, @@ -636,8 +637,9 @@ def init_new( # capture_hidden_mode=None means no override: derive from # SB.return_hidden_states / spec_info.capture_hidden_mode. + pd_hidden_capture_layer_ids = get_pd_hidden_capture_layer_ids(batch.reqs) if capture_hidden_mode is None: - if batch.return_hidden_states: + if batch.return_hidden_states or pd_hidden_capture_layer_ids: capture_hidden_mode = CaptureHiddenMode.FULL elif batch.spec_info is not None: capture_hidden_mode = getattr( @@ -708,7 +710,7 @@ def init_new( spec_algorithm=batch.spec_algorithm, capture_hidden_mode=capture_hidden_mode, return_hidden_states_before_norm=return_hidden_states_before_norm, - pd_hidden_capture_layer_ids=batch.pd_hidden_capture_layer_ids, + pd_hidden_capture_layer_ids=pd_hidden_capture_layer_ids, tbo_split_seq_index=batch.tbo_split_seq_index, # Host-side metadata top_logprobs_nums=batch.top_logprobs_nums, From 3949b6fdc8b2ee78d8676294a2d394c68fdf1562 Mon Sep 17 00:00:00 2001 From: zhangxiaolei Date: Tue, 21 Jul 2026 15:07:31 +0800 Subject: [PATCH 2/4] Move disaggregation init out of scheduler --- python/sglang/srt/disaggregation/prefill.py | 178 ++++++++++++++++++- python/sglang/srt/managers/scheduler.py | 184 +------------------- 2 files changed, 180 insertions(+), 182 deletions(-) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 8a1358f25f95..301ed3619582 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -40,6 +40,7 @@ MetadataBuffers, ReqToMetadataIdxAllocator, TransferBackend, + get_dsa_seed_metadata_dim, get_dsv4_c128_state_indices, get_kv_class, get_pd_hidden_capture_layer_ids, @@ -48,6 +49,7 @@ is_mla_backend, poll_and_all_reduce_attn_cp_tp_group, prepare_abort, + resolve_disagg_metadata_config, setup_state_kv_args, ) from sglang.srt.environ import envs @@ -122,13 +124,15 @@ def maybe_release_metadata_buffer( allocator.free(req.metadata_buffer_index) req.metadata_buffer_index = -1 indices = req.pd_hidden_src_indices - if indices and pd_hidden_pool is not None: + if indices: sender = req.disagg_kv_sender + if pd_hidden_pool is None and sender is not None: + pd_hidden_pool = sender.kv_mgr.pd_hidden_pool worker_released = ( sender is not None and sender.kv_mgr.pop_pd_hidden_request_done(sender.bootstrap_room) ) - if not worker_released: + if not worker_released and pd_hidden_pool is not None: pd_hidden_pool.free(indices) clear_pd_hidden_request_state(req) elif not indices: @@ -759,6 +763,176 @@ class SchedulerDisaggregationPrefillMixin: Mixin for Scheduler to handle disaggregation prefill """ + def init_disaggregation(self: Scheduler) -> None: + from sglang.srt.configs.model_config import is_minimax_sparse + from sglang.srt.disaggregation.decode import ( + DecodePreallocQueue, + DecodeTransferQueue, + ) + from sglang.srt.disaggregation.encode_receiver import create_mm_receiver + from sglang.srt.mem_cache import kv_cache_builder + from sglang.srt.speculative.eagle_utils import ( + get_draft_recurrent_hidden_state_spec, + ) + + self.mm_receiver = None + self.disagg_prefill_bootstrap_queue = None + self.disagg_prefill_inflight_queue = None + self.disagg_decode_prealloc_queue = None + self.disagg_decode_transfer_queue = None + + self.disaggregation_mode = DisaggregationMode( + self.server_args.disaggregation_mode + ) + self.transfer_backend = TransferBackend( + self.server_args.disaggregation_transfer_backend + ) + + # todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D? + draft_token_to_kv_pool = kv_cache_builder.get_draft_kv_pool( + draft_worker=self.draft_worker, + spec_algorithm=self.spec_algorithm, + server_args=self.server_args, + ) + + if self.spec_algorithm.carries_draft_hidden_states(): + # `draft_runner` aliases `draft_runner_list[0]` in the multi-layer + # worker, so a single accessor covers both shapes. + draft_runner = self.draft_worker.draft_worker.draft_runner + disagg_hidden_size, disagg_hidden_states_dtype = ( + get_draft_recurrent_hidden_state_spec(draft_runner) + ) + else: + disagg_hidden_size = 16 # minimal padding size for RDMA + disagg_hidden_states_dtype = torch.float32 + + disagg_metadata_config = resolve_disagg_metadata_config( + hidden_size=disagg_hidden_size, + hidden_states_dtype=disagg_hidden_states_dtype, + disaggregation_mode=self.disaggregation_mode, + transfer_backend=self.transfer_backend, + spec_algorithm=self.spec_algorithm, + model_config=self.model_config, + server_args=self.server_args, + model_runner=self.tp_worker.model_runner, + pp_rank=self.ps.pp_rank, + pp_size=self.ps.pp_size, + gpu_id=self.ps.gpu_id, + max_prefill_tokens=self.max_prefill_tokens, + ) + disagg_hidden_size = disagg_metadata_config.hidden_size + disagg_hidden_states_dtype = disagg_metadata_config.hidden_states_dtype + metadata_buffer_kwargs = disagg_metadata_config.metadata_buffer_kwargs + + # The PD metadata wire schema must match on P and D even when only D + # enables spec decoding; a seedless prefill writes the invalid sentinel. + output_dsa_topk_indices_dim = get_dsa_seed_metadata_dim( + self.model_config.hf_config + ) + + if ( + self.disaggregation_mode == DisaggregationMode.DECODE + ): # *8 headroom for MiniMax-M3; *2 for other models. + buffer_multiplier = ( + 8 if is_minimax_sparse(self.model_config.hf_config) else 2 + ) + buffer_size = (self.req_to_token_pool.size) * buffer_multiplier + self.req_to_metadata_buffer_idx_allocator = ReqToMetadataIdxAllocator( + buffer_size + ) + self.disagg_metadata_buffers = MetadataBuffers( + buffer_size, + hidden_size=disagg_hidden_size, + hidden_states_dtype=disagg_hidden_states_dtype, + custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(), + output_dsa_topk_indices_dim=output_dsa_topk_indices_dim, + **metadata_buffer_kwargs, + ) + + # The decode requests polling kv cache + self.disagg_decode_transfer_queue = DecodeTransferQueue( + gloo_group=self.attn_tp_cpu_group, + req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, + tp_rank=self.ps.tp_rank, + metadata_buffers=self.disagg_metadata_buffers, + scheduler=self, + tree_cache=self.tree_cache, + ) + + # The decode requests pending for pre-allocation + self.disagg_decode_prealloc_queue = DecodePreallocQueue( + req_to_token_pool=self.req_to_token_pool, + token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, + draft_token_to_kv_pool=draft_token_to_kv_pool, + req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, + metadata_buffers=self.disagg_metadata_buffers, + scheduler=self, + transfer_queue=self.disagg_decode_transfer_queue, + tree_cache=self.tree_cache, + gloo_group=self.attn_tp_cpu_group, + tp_rank=self.ps.tp_rank, + tp_size=self.ps.tp_size, + dp_size=self.server_args.dp_size, + gpu_id=self.ps.gpu_id, + bootstrap_port=self.server_args.disaggregation_bootstrap_port, + max_total_num_tokens=self.max_total_num_tokens, + pp_rank=self.ps.pp_rank, + num_reserved_decode_tokens=self.server_args.num_reserved_decode_tokens, + transfer_backend=self.transfer_backend, + ) + + elif self.disaggregation_mode == DisaggregationMode.PREFILL: + # *2 for the headroom. + buffer_size = self.max_running_requests * 2 + self.req_to_metadata_buffer_idx_allocator = ReqToMetadataIdxAllocator( + buffer_size + ) + self.disagg_metadata_buffers = MetadataBuffers( + buffer_size, + hidden_size=disagg_hidden_size, + hidden_states_dtype=disagg_hidden_states_dtype, + custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(), + output_dsa_topk_indices_dim=output_dsa_topk_indices_dim, + **metadata_buffer_kwargs, + ) + + self.disagg_prefill_bootstrap_queue = PrefillBootstrapQueue( + token_to_kv_pool=self.token_to_kv_pool_allocator.get_kvcache(), + draft_token_to_kv_pool=draft_token_to_kv_pool, + req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, + metadata_buffers=self.disagg_metadata_buffers, + tp_rank=self.ps.tp_rank, + tp_size=self.ps.tp_size, + gpu_id=self.ps.gpu_id, + bootstrap_port=self.server_args.disaggregation_bootstrap_port, + gloo_group=self.attn_tp_cpu_group, + max_total_num_tokens=self.max_total_num_tokens, + scheduler=self, + pp_rank=self.ps.pp_rank, + pp_size=self.ps.pp_size, + transfer_backend=self.transfer_backend, + ) + # The prefill requests that are in the middle of kv sending + self.disagg_prefill_inflight_queue: List[Req] = [] + + self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() + + # Init mm receiver for EPD disaggregation mode + if ( + self.server_args.language_only + and self.server_args.encoder_transfer_backend + in ["zmq_to_scheduler", "mooncake"] + ): + self.mm_receiver = create_mm_receiver( + self.server_args, + dtype=self.model_config.dtype, + hf_config=self.model_config.hf_config, + pp_rank=self.ps.pp_rank, + tp_rank=self.ps.tp_rank, + tp_group=self.tp_group, + scheduler=self, + ) + def maybe_prefetch_staging_for_batch(self: Scheduler, batch: ScheduleBatch) -> None: """Pre-send STAGING_REQ so decode allocates staging during GPU forward.""" kv_mgr = self.disagg_prefill_bootstrap_queue.kv_manager diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index b7d1e10c60e0..2141092bbd29 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -41,31 +41,21 @@ from sglang.kernels.ops.mamba.triton_ops import ( initialize_mamba_selective_state_update_backend, ) -from sglang.srt.configs.model_config import ModelConfig, ModelImpl, is_minimax_sparse +from sglang.srt.configs.model_config import ModelConfig, ModelImpl from sglang.srt.constrained.grammar_manager import GrammarManager from sglang.srt.debug_utils.pr_fix_toggle import maybe_revert_pr_fix -from sglang.srt.disaggregation.decode import ( - DecodePreallocQueue, - DecodeTransferQueue, - SchedulerDisaggregationDecodeMixin, -) +from sglang.srt.disaggregation.decode import SchedulerDisaggregationDecodeMixin from sglang.srt.disaggregation.decode_kvcache_offload_manager import ( DecodeKVCacheOffloadManager, ) -from sglang.srt.disaggregation.encode_receiver import create_mm_receiver from sglang.srt.disaggregation.prefill import ( - PrefillBootstrapQueue, SchedulerDisaggregationPrefillMixin, maybe_release_metadata_buffer, ) from sglang.srt.disaggregation.utils import ( DisaggregationMode, - MetadataBuffers, - ReqToMetadataIdxAllocator, TransferBackend, - get_dsa_seed_metadata_dim, prepare_abort, - resolve_disagg_metadata_config, ) from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.distributed.parallel_state import get_tp_group @@ -246,7 +236,6 @@ from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.session.session_controller import SessionController from sglang.srt.speculative.dflash_utils import validate_dflash_request -from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( DynamicGradMode, @@ -1103,165 +1092,6 @@ def init_watch_dog_memory_saver_input_blocker(self): if envs.SGLANG_LOG_GC.get(): configure_gc_logger() - def init_disaggregation(self): - self.mm_receiver = None - self.disagg_prefill_bootstrap_queue = None - self.disagg_prefill_inflight_queue = None - self.disagg_decode_prealloc_queue = None - self.disagg_decode_transfer_queue = None - - self.disaggregation_mode = DisaggregationMode( - self.server_args.disaggregation_mode - ) - self.transfer_backend = TransferBackend( - self.server_args.disaggregation_transfer_backend - ) - - # todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D? - draft_token_to_kv_pool = kv_cache_builder.get_draft_kv_pool( - draft_worker=self.draft_worker, - spec_algorithm=self.spec_algorithm, - server_args=self.server_args, - ) - - if self.spec_algorithm.carries_draft_hidden_states(): - # `draft_runner` aliases `draft_runner_list[0]` in the multi-layer - # worker, so a single accessor covers both shapes. - draft_runner = self.draft_worker.draft_worker.draft_runner - disagg_hidden_size, disagg_hidden_states_dtype = ( - get_draft_recurrent_hidden_state_spec(draft_runner) - ) - else: - disagg_hidden_size = 16 # minimal padding size for RDMA - disagg_hidden_states_dtype = torch.float32 - - disagg_metadata_config = resolve_disagg_metadata_config( - hidden_size=disagg_hidden_size, - hidden_states_dtype=disagg_hidden_states_dtype, - disaggregation_mode=self.disaggregation_mode, - transfer_backend=self.transfer_backend, - spec_algorithm=self.spec_algorithm, - model_config=self.model_config, - server_args=self.server_args, - model_runner=self.tp_worker.model_runner, - pp_rank=self.ps.pp_rank, - pp_size=self.ps.pp_size, - gpu_id=self.ps.gpu_id, - max_prefill_tokens=self.max_prefill_tokens, - ) - disagg_hidden_size = disagg_metadata_config.hidden_size - disagg_hidden_states_dtype = disagg_metadata_config.hidden_states_dtype - metadata_buffer_kwargs = disagg_metadata_config.metadata_buffer_kwargs - - # The PD metadata wire schema must match on P and D even when only D - # enables spec decoding; a seedless prefill writes the invalid sentinel. - output_dsa_topk_indices_dim = get_dsa_seed_metadata_dim( - self.model_config.hf_config - ) - - if ( - self.disaggregation_mode == DisaggregationMode.DECODE - ): # *8 headroom for MiniMax-M3; *2 for other models. - buffer_multiplier = ( - 8 if is_minimax_sparse(self.model_config.hf_config) else 2 - ) - buffer_size = (self.req_to_token_pool.size) * buffer_multiplier - self.req_to_metadata_buffer_idx_allocator = ReqToMetadataIdxAllocator( - buffer_size - ) - self.disagg_metadata_buffers = MetadataBuffers( - buffer_size, - hidden_size=disagg_hidden_size, - hidden_states_dtype=disagg_hidden_states_dtype, - custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(), - output_dsa_topk_indices_dim=output_dsa_topk_indices_dim, - **metadata_buffer_kwargs, - ) - - # The decode requests polling kv cache - self.disagg_decode_transfer_queue = DecodeTransferQueue( - gloo_group=self.attn_tp_cpu_group, - req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, - tp_rank=self.ps.tp_rank, - metadata_buffers=self.disagg_metadata_buffers, - scheduler=self, - tree_cache=self.tree_cache, - ) - - # The decode requests pending for pre-allocation - self.disagg_decode_prealloc_queue = DecodePreallocQueue( - req_to_token_pool=self.req_to_token_pool, - token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, - draft_token_to_kv_pool=draft_token_to_kv_pool, - req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, - metadata_buffers=self.disagg_metadata_buffers, - scheduler=self, - transfer_queue=self.disagg_decode_transfer_queue, - tree_cache=self.tree_cache, - gloo_group=self.attn_tp_cpu_group, - tp_rank=self.ps.tp_rank, - tp_size=self.ps.tp_size, - dp_size=self.server_args.dp_size, - gpu_id=self.ps.gpu_id, - bootstrap_port=self.server_args.disaggregation_bootstrap_port, - max_total_num_tokens=self.max_total_num_tokens, - pp_rank=self.ps.pp_rank, - num_reserved_decode_tokens=self.server_args.num_reserved_decode_tokens, - transfer_backend=self.transfer_backend, - ) - - elif self.disaggregation_mode == DisaggregationMode.PREFILL: - # *2 for the headroom. - buffer_size = self.max_running_requests * 2 - self.req_to_metadata_buffer_idx_allocator = ReqToMetadataIdxAllocator( - buffer_size - ) - self.disagg_metadata_buffers = MetadataBuffers( - buffer_size, - hidden_size=disagg_hidden_size, - hidden_states_dtype=disagg_hidden_states_dtype, - custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(), - output_dsa_topk_indices_dim=output_dsa_topk_indices_dim, - **metadata_buffer_kwargs, - ) - - self.disagg_prefill_bootstrap_queue = PrefillBootstrapQueue( - token_to_kv_pool=self.token_to_kv_pool_allocator.get_kvcache(), - draft_token_to_kv_pool=draft_token_to_kv_pool, - req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, - metadata_buffers=self.disagg_metadata_buffers, - tp_rank=self.ps.tp_rank, - tp_size=self.ps.tp_size, - gpu_id=self.ps.gpu_id, - bootstrap_port=self.server_args.disaggregation_bootstrap_port, - gloo_group=self.attn_tp_cpu_group, - max_total_num_tokens=self.max_total_num_tokens, - scheduler=self, - pp_rank=self.ps.pp_rank, - pp_size=self.ps.pp_size, - transfer_backend=self.transfer_backend, - ) - # The prefill requests that are in the middle of kv sending - self.disagg_prefill_inflight_queue: List[Req] = [] - - self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() - - # Init mm receiver for EPD disaggregation mode - if ( - self.server_args.language_only - and self.server_args.encoder_transfer_backend - in ["zmq_to_scheduler", "mooncake"] - ): - self.mm_receiver = create_mm_receiver( - self.server_args, - dtype=self.model_config.dtype, - hf_config=self.model_config.hf_config, - pp_rank=self.ps.pp_rank, - tp_rank=self.ps.tp_rank, - tp_group=self.tp_group, - scheduler=self, - ) - def init_overlap(self): self.device_module = torch.get_device_module(self.device) @@ -2582,9 +2412,7 @@ def process_pending_chunked_abort(self) -> None: if self.disaggregation_mode == DisaggregationMode.PREFILL: req.disagg_kv_sender.abort() maybe_release_metadata_buffer( - req, - self.req_to_metadata_buffer_idx_allocator, - getattr(self.disagg_metadata_buffers, "pd_hidden_pool", None), + req, self.req_to_metadata_buffer_idx_allocator ) req.pending_bootstrap = False if self.enable_hicache_storage: @@ -3992,11 +3820,7 @@ def abort_request(self, recv_req: AbortReq): if self.disaggregation_mode == DisaggregationMode.PREFILL: bootstrap_pending = req.pending_bootstrap maybe_release_metadata_buffer( - req, - self.req_to_metadata_buffer_idx_allocator, - getattr( - self.disagg_metadata_buffers, "pd_hidden_pool", None - ), + req, self.req_to_metadata_buffer_idx_allocator ) if ( bootstrap_pending From 3d2dbb0d2f1f809473fd4feac3eec892b7b40e86 Mon Sep 17 00:00:00 2001 From: zhangxiaolei Date: Tue, 21 Jul 2026 15:28:14 +0800 Subject: [PATCH 3/4] Move PD hidden request state out of Req --- python/sglang/srt/disaggregation/prefill.py | 145 +++++++++--------- python/sglang/srt/disaggregation/utils.py | 28 +++- python/sglang/srt/managers/schedule_batch.py | 10 -- .../sglang/srt/managers/scheduler_pp_mixin.py | 9 +- 4 files changed, 103 insertions(+), 89 deletions(-) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 301ed3619582..c63ba0aa1143 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -44,6 +44,7 @@ get_dsv4_c128_state_indices, get_kv_class, get_pd_hidden_capture_layer_ids, + get_pd_hidden_req_state as pd_hidden_state, is_aborted, is_dsv4_c128_online_enabled, is_mla_backend, @@ -94,16 +95,16 @@ def should_force_retry(req: Req) -> bool: def clear_pd_hidden_request_state(req: Req) -> None: - req.pd_hidden_meta = None - req.pd_hidden_src_indices = None - req.pd_hidden_dst_indices = None - req.pd_hidden_written = None - req.pd_hidden_capture_layer_ids = None - req.pd_hidden_current_src_indices = None - req.pd_hidden_current_start = None - req.pd_hidden_current_row_len = 0 - req.pd_hidden_current_is_last = False - req.pd_hidden_owner_direct_sent = False + pd_hidden_state(req).meta = None + pd_hidden_state(req).src_indices = None + pd_hidden_state(req).dst_indices = None + pd_hidden_state(req).written = None + pd_hidden_state(req).capture_layer_ids = None + pd_hidden_state(req).current_src_indices = None + pd_hidden_state(req).current_start = None + pd_hidden_state(req).current_row_len = 0 + pd_hidden_state(req).current_is_last = False + pd_hidden_state(req).owner_direct_sent = False def maybe_release_metadata_buffer( @@ -123,7 +124,7 @@ def maybe_release_metadata_buffer( if req.metadata_buffer_index >= 0: allocator.free(req.metadata_buffer_index) req.metadata_buffer_index = -1 - indices = req.pd_hidden_src_indices + indices = pd_hidden_state(req).src_indices if indices: sender = req.disagg_kv_sender if pd_hidden_pool is None and sender is not None: @@ -143,7 +144,7 @@ def maybe_release_pd_hidden_rows(req: Req, pd_hidden_pool) -> None: """Release source hidden rows once the local RDMA transfer is complete.""" if pd_hidden_pool is None: return - indices = req.pd_hidden_src_indices + indices = pd_hidden_state(req).src_indices if indices: pd_hidden_pool.free(indices) clear_pd_hidden_request_state(req) @@ -153,7 +154,7 @@ def maybe_release_pd_hidden_rows_on_hidden_done( req: Req, pd_hidden_pool ) -> bool: """Release source hidden rows after PD_HIDDEN finishes, before KV success.""" - indices = req.pd_hidden_src_indices + indices = pd_hidden_state(req).src_indices if not indices or pd_hidden_pool is None: return False sender = req.disagg_kv_sender @@ -452,7 +453,7 @@ def _probe_bootstrap_ready( return (metadata_cost, 0), None hidden_cost = 0 if plan.streaming_hidden else plan.source_window_rows - if req.pd_hidden_src_indices is not None: + if pd_hidden_state(req).src_indices is not None: hidden_cost = 0 if hidden_cost > hidden_row_credits: now = time.monotonic() @@ -494,7 +495,7 @@ def _is_pd_hidden_credit_blocked( else [int(x) for x in dspark_meta.get("target_layer_ids", [])] ) ) - if not local_layer_ids or req.pd_hidden_src_indices: + if not local_layer_ids or pd_hidden_state(req).src_indices: return False pool = getattr(self.metadata_buffers, "pd_hidden_pool", None) @@ -549,11 +550,11 @@ def _finalize_pd_hidden_bootstrap( assert plan is not None if not plan.local_layer_ids: - req.pd_hidden_meta = dict(dspark_meta) - req.pd_hidden_src_indices = [] - req.pd_hidden_dst_indices = [] - req.pd_hidden_written = [] - req.pd_hidden_owner_direct_sent = False + pd_hidden_state(req).meta = dict(dspark_meta) + pd_hidden_state(req).src_indices = [] + pd_hidden_state(req).dst_indices = [] + pd_hidden_state(req).written = [] + pd_hidden_state(req).owner_direct_sent = False return True src_indices = ( @@ -573,14 +574,14 @@ def _finalize_pd_hidden_bootstrap( self._abort_pd_hidden_bootstrap(req, message) return False - req.pd_hidden_capture_layer_ids = [int(x) for x in plan.local_layer_ids] - req.pd_hidden_meta = dict(dspark_meta) - req.pd_hidden_src_indices = src_indices - req.pd_hidden_dst_indices = plan.dst_indices - req.pd_hidden_written = ( + pd_hidden_state(req).capture_layer_ids = [int(x) for x in plan.local_layer_ids] + pd_hidden_state(req).meta = dict(dspark_meta) + pd_hidden_state(req).src_indices = src_indices + pd_hidden_state(req).dst_indices = plan.dst_indices + pd_hidden_state(req).written = ( None if plan.streaming_hidden else [False] * plan.hidden_len ) - req.pd_hidden_owner_direct_sent = False + pd_hidden_state(req).owner_direct_sent = False return True def add(self, req: Req, num_kv_heads: int) -> None: @@ -1126,7 +1127,7 @@ def _extract_pd_hidden_states_from_result( def _build_pd_hidden_only_state_indices( self: Scheduler, req: Req ) -> Optional[List]: - current_indices = req.pd_hidden_current_src_indices + current_indices = pd_hidden_state(req).current_src_indices if current_indices is None: return None @@ -1142,9 +1143,9 @@ def _build_pd_hidden_only_state_indices( return state_indices def _send_pd_hidden_only_chunk(self: Scheduler, req: Req) -> bool: - current_indices = req.pd_hidden_current_src_indices - current_start = req.pd_hidden_current_start - current_rows = int(req.pd_hidden_current_row_len or 0) + current_indices = pd_hidden_state(req).current_src_indices + current_start = pd_hidden_state(req).current_start + current_rows = int(pd_hidden_state(req).current_row_len or 0) if current_indices is None or current_start is None or current_rows <= 0: return False @@ -1153,7 +1154,7 @@ def _send_pd_hidden_only_chunk(self: Scheduler, req: Req) -> bool: return False streaming_hidden = bool( - (req.pd_hidden_meta or {}).get("streaming_hidden", False) + (pd_hidden_state(req).meta or {}).get("streaming_hidden", False) ) if req.disagg_kv_sender is not None: source_event = self.device_module.Event() @@ -1162,18 +1163,18 @@ def _send_pd_hidden_only_chunk(self: Scheduler, req: Req) -> bool: req.disagg_kv_sender.set_pd_hidden_chunk_meta( int(current_start), int(current_rows), - bool(req.pd_hidden_current_is_last), - current_indices if streaming_hidden else req.pd_hidden_src_indices, + bool(pd_hidden_state(req).current_is_last), + current_indices if streaming_hidden else pd_hidden_state(req).src_indices, ) req.disagg_kv_sender.send(np.asarray([], dtype=np.int32), state_indices) if streaming_hidden: - req.pd_hidden_src_indices = None - req.pd_hidden_current_src_indices = None - req.pd_hidden_current_start = None - req.pd_hidden_current_row_len = 0 - req.pd_hidden_current_is_last = False - req.pd_hidden_owner_direct_sent = True + pd_hidden_state(req).src_indices = None + pd_hidden_state(req).current_src_indices = None + pd_hidden_state(req).current_start = None + pd_hidden_state(req).current_row_len = 0 + pd_hidden_state(req).current_is_last = False + pd_hidden_state(req).owner_direct_sent = True return True def _write_pd_hidden_rows_for_batch( @@ -1190,12 +1191,12 @@ def _write_pd_hidden_rows_for_batch( for req in batch.reqs if ( ( - req.pd_hidden_src_indices - or req.pd_hidden_capture_layer_ids + pd_hidden_state(req).src_indices + or pd_hidden_state(req).capture_layer_ids ) and ( send_owner_direct - or not req.pd_hidden_owner_direct_sent + or not pd_hidden_state(req).owner_direct_sent ) ) ] @@ -1203,8 +1204,8 @@ def _write_pd_hidden_rows_for_batch( reqs = [ ( req.rid, - req.pd_hidden_capture_layer_ids, - bool(req.pd_hidden_src_indices), + pd_hidden_state(req).capture_layer_ids, + bool(pd_hidden_state(req).src_indices), ) for req in needs_pd_hidden_reqs ] @@ -1239,11 +1240,11 @@ def _write_pd_hidden_rows_for_batch( req_hidden = hidden_states[hidden_offset : hidden_offset + extend_len] hidden_offset += extend_len - meta = req.pd_hidden_meta or {} + meta = pd_hidden_state(req).meta or {} streaming_hidden = bool(meta.get("streaming_hidden", False)) - if not send_owner_direct and req.pd_hidden_owner_direct_sent: + if not send_owner_direct and pd_hidden_state(req).owner_direct_sent: continue - src_indices = req.pd_hidden_src_indices + src_indices = pd_hidden_state(req).src_indices if not src_indices and not streaming_hidden: continue @@ -1291,8 +1292,8 @@ def _write_pd_hidden_rows_for_batch( else: rows = local_end - local_start write_indices = src_indices[local_start:local_end] - prev_current_start = req.pd_hidden_current_start - prev_current_row_len = int(req.pd_hidden_current_row_len or 0) + prev_current_start = pd_hidden_state(req).current_start + prev_current_row_len = int(pd_hidden_state(req).current_row_len or 0) if ( prev_current_start is not None and prev_current_row_len > 0 @@ -1322,16 +1323,16 @@ def _write_pd_hidden_rows_for_batch( "only after the matching hidden chunk ACK.", ) continue - req.pd_hidden_src_indices = write_indices + pd_hidden_state(req).src_indices = write_indices pool.write( write_indices, req_hidden_to_write[chunk_local_start:chunk_local_end], ) - req.pd_hidden_current_start = write_start - req.pd_hidden_current_row_len = rows - req.pd_hidden_current_src_indices = write_indices - req.pd_hidden_current_is_last = write_end >= hidden_start + hidden_len - written = req.pd_hidden_written + pd_hidden_state(req).current_start = write_start + pd_hidden_state(req).current_row_len = rows + pd_hidden_state(req).current_src_indices = write_indices + pd_hidden_state(req).current_is_last = write_end >= hidden_start + hidden_len + written = pd_hidden_state(req).written if written is not None: written[local_start:local_end] = [True] * rows if send_owner_direct: @@ -1345,7 +1346,7 @@ def send_dspark_owner_direct_hidden_for_batch( capture_reqs = [ req for req in batch.reqs - if req.pd_hidden_capture_layer_ids + if pd_hidden_state(req).capture_layer_ids ] if not capture_reqs: return False @@ -1353,7 +1354,7 @@ def send_dspark_owner_direct_hidden_for_batch( return False if not all( bool( - (req.pd_hidden_meta or {}).get("streaming_hidden", False) + (pd_hidden_state(req).meta or {}).get("streaming_hidden", False) ) for req in capture_reqs ): @@ -1871,15 +1872,15 @@ def send_kv_chunk( ) return True - current_pd_hidden_src_indices = req.pd_hidden_current_src_indices - current_pd_hidden_start = req.pd_hidden_current_start - current_pd_hidden_row_len = int(req.pd_hidden_current_row_len or 0) + current_pd_hidden_src_indices = pd_hidden_state(req).current_src_indices + current_pd_hidden_start = pd_hidden_state(req).current_start + current_pd_hidden_row_len = int(pd_hidden_state(req).current_row_len or 0) has_current_pd_hidden = ( current_pd_hidden_src_indices is not None and current_pd_hidden_row_len > 0 ) streaming_pd_hidden = bool( - (req.pd_hidden_meta or {}).get("streaming_hidden", False) + (pd_hidden_state(req).meta or {}).get("streaming_hidden", False) ) state_indices: Optional[List] = None @@ -1953,10 +1954,10 @@ def _c128_state_payload(): ) def _pd_hidden_payload(): - if req.pd_hidden_owner_direct_sent: + if pd_hidden_state(req).owner_direct_sent: return [] - src_indices = req.pd_hidden_src_indices - if src_indices is None and req.pd_hidden_capture_layer_ids: + src_indices = pd_hidden_state(req).src_indices + if src_indices is None and pd_hidden_state(req).capture_layer_ids: raise RuntimeError( "PD hidden row pool was not materialized before transfer: " f"rid={req.rid}" @@ -1965,7 +1966,7 @@ def _pd_hidden_payload(): return np.asarray(current_pd_hidden_src_indices, dtype=np.int32) if not src_indices: return [] - written = req.pd_hidden_written + written = pd_hidden_state(req).written if written is not None and not all(written): missing = [i for i, ok in enumerate(written) if not ok][:8] raise RuntimeError( @@ -2014,18 +2015,18 @@ def _pd_hidden_payload(): req.disagg_kv_sender.set_pd_hidden_chunk_meta( int(current_pd_hidden_start), int(current_pd_hidden_row_len), - bool(req.pd_hidden_current_is_last), + bool(pd_hidden_state(req).current_is_last), current_pd_hidden_src_indices if streaming_pd_hidden - else req.pd_hidden_src_indices, + else pd_hidden_state(req).src_indices, ) req.disagg_kv_sender.send(page_indices, state_indices) if has_current_pd_hidden and streaming_pd_hidden: - req.pd_hidden_src_indices = None - req.pd_hidden_current_src_indices = None - req.pd_hidden_current_start = None - req.pd_hidden_current_row_len = 0 - req.pd_hidden_current_is_last = False + pd_hidden_state(req).src_indices = None + pd_hidden_state(req).current_src_indices = None + pd_hidden_state(req).current_start = None + pd_hidden_state(req).current_row_len = 0 + pd_hidden_state(req).current_is_last = False req.start_send_idx = end_idx return True diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 8f039aef184b..a939d9f9fe28 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -4,6 +4,7 @@ import random import logging import threading +import weakref from collections import deque from contextlib import nullcontext from enum import Enum @@ -83,12 +84,37 @@ def get_dsv4_c128_state_indices( def get_pd_hidden_capture_layer_ids(reqs: List["Req"]) -> Optional[List[int]]: """Return the per-batch PD hidden capture layers requested by any req.""" for req in reqs: - layer_ids = getattr(req, "pd_hidden_capture_layer_ids", None) + layer_ids = get_pd_hidden_req_state(req).capture_layer_ids if layer_ids: return [int(x) for x in layer_ids] return None +class PDHiddenReqState: + def __init__(self): + self.meta: Optional[dict] = None + self.src_indices: Optional[List[int]] = None + self.dst_indices: Optional[List[int]] = None + self.written: Optional[List[bool]] = None + self.capture_layer_ids: Optional[List[int]] = None + self.current_src_indices: Optional[List[int]] = None + self.current_start: Optional[int] = None + self.current_row_len: int = 0 + self.current_is_last: bool = False + self.owner_direct_sent: bool = False + + +_pd_hidden_req_states = weakref.WeakKeyDictionary() + + +def get_pd_hidden_req_state(req: "Req") -> PDHiddenReqState: + state = _pd_hidden_req_states.get(req) + if state is None: + state = PDHiddenReqState() + _pd_hidden_req_states[req] = state + return state + + class DisaggregationMode(Enum): NULL = "null" PREFILL = "prefill" diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 40d90adfe5ac..9f0ec279e83e 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1050,16 +1050,6 @@ def __init__( self.pd_rebootstrap_forced_output_id: Optional[int] = None self.skip_radix_cache_insert = bootstrap_host == FAKE_BOOTSTRAP_HOST self.disagg_kv_sender: Optional[BaseKVSender] = None - self.pd_hidden_meta: Optional[dict] = None - self.pd_hidden_src_indices: Optional[List[int]] = None - self.pd_hidden_dst_indices: Optional[List[int]] = None - self.pd_hidden_written: Optional[List[bool]] = None - self.pd_hidden_capture_layer_ids: Optional[List[int]] = None - self.pd_hidden_current_src_indices: Optional[List[int]] = None - self.pd_hidden_current_start: Optional[int] = None - self.pd_hidden_current_row_len: int = 0 - self.pd_hidden_current_is_last: bool = False - self.pd_hidden_owner_direct_sent: bool = False self.routed_dp_rank: Optional[int] = routed_dp_rank self.disagg_prefill_dp_rank: Optional[int] = disagg_prefill_dp_rank diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index f4dab0962e8e..f58aa68bf685 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -16,6 +16,7 @@ from sglang.srt.disaggregation.base.conn import KVPoll from sglang.srt.disaggregation.utils import ( get_pd_hidden_capture_layer_ids, + get_pd_hidden_req_state as pd_hidden_state, poll_and_all_reduce_attn_cp_tp_group, ) from sglang.srt.distributed.parallel_state import P2PWork @@ -1036,18 +1037,14 @@ def _pp_should_owner_direct_pd_hidden( capture_reqs = [ req for req in batch.reqs - if getattr(req, "pd_hidden_capture_layer_ids", None) + if pd_hidden_state(req).capture_layer_ids ] if not capture_reqs: return False if any(req.pending_bootstrap for req in capture_reqs): return False return all( - bool( - (getattr(req, "pd_hidden_meta", None) or {}).get( - "streaming_hidden", False - ) - ) + bool((pd_hidden_state(req).meta or {}).get("streaming_hidden", False)) for req in capture_reqs ) From 4772b91505a92915e24960d541bf3e2fc73faeac Mon Sep 17 00:00:00 2001 From: zhangxiaolei Date: Tue, 21 Jul 2026 15:39:02 +0800 Subject: [PATCH 4/4] Break PD hidden state import cycle --- .../sglang/srt/disaggregation/hidden_state.py | 41 +++++++++++++++++++ python/sglang/srt/disaggregation/prefill.py | 6 ++- python/sglang/srt/disaggregation/utils.py | 35 ---------------- .../sglang/srt/managers/scheduler_pp_mixin.py | 4 +- .../srt/model_executor/forward_batch_info.py | 2 +- 5 files changed, 48 insertions(+), 40 deletions(-) create mode 100644 python/sglang/srt/disaggregation/hidden_state.py diff --git a/python/sglang/srt/disaggregation/hidden_state.py b/python/sglang/srt/disaggregation/hidden_state.py new file mode 100644 index 000000000000..5f5e7f6ff9e8 --- /dev/null +++ b/python/sglang/srt/disaggregation/hidden_state.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +import weakref +from typing import TYPE_CHECKING, List, Optional + +if TYPE_CHECKING: + from sglang.srt.managers.schedule_batch import Req + + +class PDHiddenReqState: + def __init__(self): + self.meta: Optional[dict] = None + self.src_indices: Optional[List[int]] = None + self.dst_indices: Optional[List[int]] = None + self.written: Optional[List[bool]] = None + self.capture_layer_ids: Optional[List[int]] = None + self.current_src_indices: Optional[List[int]] = None + self.current_start: Optional[int] = None + self.current_row_len: int = 0 + self.current_is_last: bool = False + self.owner_direct_sent: bool = False + + +_pd_hidden_req_states = weakref.WeakKeyDictionary() + + +def get_pd_hidden_req_state(req: "Req") -> PDHiddenReqState: + state = _pd_hidden_req_states.get(req) + if state is None: + state = PDHiddenReqState() + _pd_hidden_req_states[req] = state + return state + + +def get_pd_hidden_capture_layer_ids(reqs: List["Req"]) -> Optional[List[int]]: + """Return the per-batch PD hidden capture layers requested by any req.""" + for req in reqs: + layer_ids = get_pd_hidden_req_state(req).capture_layer_ids + if layer_ids: + return [int(x) for x in layer_ids] + return None diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index c63ba0aa1143..1229d5300149 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -33,6 +33,10 @@ from sglang.srt.disaggregation.base import KVPoll from sglang.srt.disaggregation.base.conn import StateType from sglang.srt.disaggregation.common.conn import CommonKVManager +from sglang.srt.disaggregation.hidden_state import ( + get_pd_hidden_capture_layer_ids, + get_pd_hidden_req_state as pd_hidden_state, +) from sglang.srt.disaggregation.utils import ( FAKE_BOOTSTRAP_HOST, DisaggregationMode, @@ -43,8 +47,6 @@ get_dsa_seed_metadata_dim, get_dsv4_c128_state_indices, get_kv_class, - get_pd_hidden_capture_layer_ids, - get_pd_hidden_req_state as pd_hidden_state, is_aborted, is_dsv4_c128_online_enabled, is_mla_backend, diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index a939d9f9fe28..379a7a0e7032 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -4,7 +4,6 @@ import random import logging import threading -import weakref from collections import deque from contextlib import nullcontext from enum import Enum @@ -81,40 +80,6 @@ def get_dsv4_c128_state_indices( return np.array([page], dtype=np.int32) -def get_pd_hidden_capture_layer_ids(reqs: List["Req"]) -> Optional[List[int]]: - """Return the per-batch PD hidden capture layers requested by any req.""" - for req in reqs: - layer_ids = get_pd_hidden_req_state(req).capture_layer_ids - if layer_ids: - return [int(x) for x in layer_ids] - return None - - -class PDHiddenReqState: - def __init__(self): - self.meta: Optional[dict] = None - self.src_indices: Optional[List[int]] = None - self.dst_indices: Optional[List[int]] = None - self.written: Optional[List[bool]] = None - self.capture_layer_ids: Optional[List[int]] = None - self.current_src_indices: Optional[List[int]] = None - self.current_start: Optional[int] = None - self.current_row_len: int = 0 - self.current_is_last: bool = False - self.owner_direct_sent: bool = False - - -_pd_hidden_req_states = weakref.WeakKeyDictionary() - - -def get_pd_hidden_req_state(req: "Req") -> PDHiddenReqState: - state = _pd_hidden_req_states.get(req) - if state is None: - state = PDHiddenReqState() - _pd_hidden_req_states[req] = state - return state - - class DisaggregationMode(Enum): NULL = "null" PREFILL = "prefill" diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index f58aa68bf685..2b61bb6b4fc6 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -14,11 +14,11 @@ from tqdm import tqdm from sglang.srt.disaggregation.base.conn import KVPoll -from sglang.srt.disaggregation.utils import ( +from sglang.srt.disaggregation.hidden_state import ( get_pd_hidden_capture_layer_ids, get_pd_hidden_req_state as pd_hidden_state, - poll_and_all_reduce_attn_cp_tp_group, ) +from sglang.srt.disaggregation.utils import poll_and_all_reduce_attn_cp_tp_group from sglang.srt.distributed.parallel_state import P2PWork from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import ( diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index c39a94b3f0e8..e5790ffb45e3 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -42,7 +42,7 @@ from sglang.srt.kv_canary.req_to_expected_token_ids_manager import ( compute_req_all_ids_info, ) -from sglang.srt.disaggregation.utils import get_pd_hidden_capture_layer_ids +from sglang.srt.disaggregation.hidden_state import get_pd_hidden_capture_layer_ids from sglang.srt.layers.dp_attention import ( DpPaddingMode, set_dp_buffer_len,