From 38c9845005f873b99243c4a80d3815f636840e36 Mon Sep 17 00:00:00 2001 From: zhangxiaolei Date: Mon, 20 Jul 2026 22:38:27 +0800 Subject: [PATCH] Rename DSpark hidden protocol to PD hidden --- python/sglang/srt/disaggregation/base/conn.py | 2 +- .../sglang/srt/disaggregation/common/conn.py | 20 +- .../sglang/srt/disaggregation/common/utils.py | 24 +- python/sglang/srt/disaggregation/decode.py | 152 ++-- .../srt/disaggregation/mooncake/conn.py | 652 +++++++++--------- python/sglang/srt/disaggregation/prefill.py | 284 ++++---- python/sglang/srt/disaggregation/utils.py | 17 +- python/sglang/srt/environ.py | 2 +- python/sglang/srt/managers/schedule_batch.py | 8 +- .../sglang/srt/managers/scheduler_pp_mixin.py | 30 +- .../srt/model_executor/forward_batch_info.py | 4 +- python/sglang/srt/models/deepseek_v4.py | 38 +- .../dspark_disaggregation.py | 14 +- ...idden_state.py => test_pd_hidden_state.py} | 38 +- 14 files changed, 636 insertions(+), 649 deletions(-) rename test/registered/unit/disaggregation/{test_dspark_hidden_state.py => test_pd_hidden_state.py} (83%) diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 8e6b342c5581..1a565f7c2b17 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -25,7 +25,7 @@ class StateType(str, enum.Enum): # DeepSeek-V4 online C128 request-scoped state. C128_STATE = "c128_state" # Target aux hidden rows used to bootstrap decode-side draft KV. - DSPARK_HIDDEN = "dspark_hidden" + PD_HIDDEN = "pd_hidden" @dataclasses.dataclass diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index af63990edc4c..c712481b1320 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -200,7 +200,7 @@ def __init__( self.register_to_bootstrap() self.transfer_infos = {} self.req_to_decode_prefix_len: Dict[int, int] = {} - self.req_to_dspark_hidden_meta: Dict[int, dict] = {} + self.req_to_pd_hidden_meta: Dict[int, dict] = {} self.decode_kv_args_table = {} self.pp_group = get_pp_group() # If a timeout happens on the prefill side, it means prefill instances @@ -245,10 +245,10 @@ def __init__( f"Unsupported DisaggregationMode: {self.disaggregation_mode}" ) - def supports_dspark_hidden_streaming(self) -> bool: + def supports_pd_hidden_streaming(self) -> bool: return False - def mark_dspark_hidden_request_done( + def mark_pd_hidden_request_done( self, bootstrap_room: int, state_indices: Optional[List] = None, @@ -261,22 +261,22 @@ def mark_dspark_hidden_request_done( del bootstrap_room, state_indices return None - def pop_dspark_hidden_request_done(self, bootstrap_room: int) -> bool: + def pop_pd_hidden_request_done(self, bootstrap_room: int) -> bool: """Consume a hidden-request-done event for early source-window release.""" del bootstrap_room return False # Backward-compatible aliases for backend-specific implementations that have # not yet migrated to the request-level naming. - def mark_dspark_hidden_done( + def mark_pd_hidden_done( self, bootstrap_room: int, state_indices: Optional[List] = None, ) -> None: - self.mark_dspark_hidden_request_done(bootstrap_room, state_indices) + self.mark_pd_hidden_request_done(bootstrap_room, state_indices) - def pop_dspark_hidden_done(self, bootstrap_room: int) -> bool: - return self.pop_dspark_hidden_request_done(bootstrap_room) + def pop_pd_hidden_done(self, bootstrap_room: int) -> bool: + return self.pop_pd_hidden_request_done(bootstrap_room) def check_status(self, bootstrap_room: int) -> KVPoll: return self.request_status[bootstrap_room] @@ -1182,8 +1182,8 @@ def clear(self) -> None: self.kv_mgr.request_status.pop(self.bootstrap_room, None) if hasattr(self.kv_mgr, "req_to_decode_prefix_len"): self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, None) - if hasattr(self.kv_mgr, "req_to_dspark_hidden_meta"): - self.kv_mgr.req_to_dspark_hidden_meta.pop(self.bootstrap_room, None) + if hasattr(self.kv_mgr, "req_to_pd_hidden_meta"): + self.kv_mgr.req_to_pd_hidden_meta.pop(self.bootstrap_room, None) if hasattr(self.kv_mgr, "transfer_infos"): self.kv_mgr.transfer_infos.pop(self.bootstrap_room, None) diff --git a/python/sglang/srt/disaggregation/common/utils.py b/python/sglang/srt/disaggregation/common/utils.py index 4c683efed403..0a25c630afd9 100644 --- a/python/sglang/srt/disaggregation/common/utils.py +++ b/python/sglang/srt/disaggregation/common/utils.py @@ -26,16 +26,16 @@ class TransferKVChunk: state_indices: Optional[List] chunk_id: Optional[int] = None kv_sent: bool = False - dspark_hidden_packet_idx: int = 0 - dspark_hidden_sent: bool = False - dspark_hidden_ready_sent: bool = False - dspark_hidden_ack_ready: bool = False - dspark_hidden_ack_expected_count: int = 0 - dspark_hidden_ack_timed_out: bool = False - dspark_hidden_start: Optional[int] = None - dspark_hidden_row_len: int = 0 - dspark_hidden_is_last_chunk: bool = False - dspark_hidden_release_indices: Optional[List[int]] = None + pd_hidden_packet_idx: int = 0 + pd_hidden_sent: bool = False + pd_hidden_ready_sent: bool = False + pd_hidden_ack_ready: bool = False + pd_hidden_ack_expected_count: int = 0 + pd_hidden_ack_timed_out: bool = False + pd_hidden_start: Optional[int] = None + pd_hidden_row_len: int = 0 + pd_hidden_is_last_chunk: bool = False + pd_hidden_release_indices: Optional[List[int]] = None enqueue_time: float = 0.0 source_event: Optional[Any] = None trace_ctx: Union[TraceReqContext, TraceNullContext] = dataclasses.field( @@ -146,10 +146,6 @@ def accept_chunk( return "accepted" -DSparkHiddenChunk = PDHiddenChunk -DSparkHiddenRequestState = PDHiddenRequestState - - def pack_list_of_buffers(buffers: List[bytes]) -> bytes: if not buffers: return b"" diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index ba409d6459e5..c738928f47a7 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -269,11 +269,11 @@ class DecodeRequest: waiting_for_input: bool = False metadata_buffer_index: int = -1 is_rebootstrap: bool = False - dspark_hidden_dst_indices: Optional[List[int]] = None - dspark_hidden_dst_indices_by_pp: Optional[Dict[int, List[int]]] = None - dspark_hidden_pp_slices: Optional[Dict[int, dict]] = None - dspark_hidden_start: int = 0 - dspark_hidden_state: PDHiddenRequestState = field( + pd_hidden_dst_indices: Optional[List[int]] = None + pd_hidden_dst_indices_by_pp: Optional[Dict[int, List[int]]] = None + pd_hidden_pp_slices: Optional[Dict[int, dict]] = None + pd_hidden_start: int = 0 + pd_hidden_state: PDHiddenRequestState = field( default_factory=PDHiddenRequestState.disabled ) @@ -347,7 +347,7 @@ def __init__( self._max_ensure_retries: int = 15 # scheduling cycles self._ensure_last_attempt_time: Dict[str, float] = {} self._ensure_retry_interval: float = 1.0 # seconds - self._last_dspark_hidden_recv_credit_warning_time = 0.0 + self._last_pd_hidden_recv_credit_warning_time = 0.0 # Retracted requests staged for rebootstrap while generation is paused. # Enqueued into ``self.queue`` only on ``continue_generation`` so the # prefix KV is recomputed under the post-retract (updated) weights. @@ -928,7 +928,7 @@ def pop_preallocated( ) decode_req.kv_receiver.clear() decode_req.kv_receiver = None - self._release_dspark_hidden_rows(decode_req) + self._release_pd_hidden_rows(decode_req) failed_reqs.append(decode_req) indices_to_remove.add(i) @@ -1057,19 +1057,19 @@ def pop_preallocated( self.tree_cache.dec_lock_ref(decode_req.req.last_node) break - dspark_hidden_dst_indices = None - dspark_hidden_dst_indices_by_pp = None - dspark_hidden_pp_slices = None - dspark_hidden_start = total_prefix_len - dspark_hidden_len = origin_input_len - total_prefix_len + pd_hidden_dst_indices = None + pd_hidden_dst_indices_by_pp = None + pd_hidden_pp_slices = None + pd_hidden_start = total_prefix_len + pd_hidden_len = origin_input_len - total_prefix_len state_types = self.kv_manager.kv_args.state_types if ( self.scheduler.spec_algorithm.is_dspark() and not _is_fake_transfer( decode_req.req, self.scheduler.server_args ) - and StateType.DSPARK_HIDDEN in state_types - and dspark_hidden_len > 0 + and StateType.PD_HIDDEN in state_types + and pd_hidden_len > 0 ): dspark_pool = getattr(self.metadata_buffers, "pd_hidden_pool", None) if dspark_pool is None: @@ -1205,11 +1205,11 @@ def pop_preallocated( indices_to_remove.add(i) continue - dspark_hidden_streaming = ( - self.kv_manager.supports_dspark_hidden_streaming() + pd_hidden_streaming = ( + self.kv_manager.supports_pd_hidden_streaming() and hasattr(self.scheduler.draft_worker, "inject_pd_hidden_chunk") ) - if not dspark_hidden_streaming: + if not pd_hidden_streaming: message = ( "PD hidden transfer requires streaming chunk injection support." ) @@ -1227,13 +1227,13 @@ def pop_preallocated( failed_reqs.append(decode_req) indices_to_remove.add(i) continue - dspark_hidden_window_rows = min(dspark_hidden_len, dspark_pool.size) - if dspark_hidden_window_rows <= 0: + pd_hidden_window_rows = min(pd_hidden_len, dspark_pool.size) + if pd_hidden_window_rows <= 0: message = ( "PD decode hidden receive pool has no streaming rows: " - f"rid={decode_req.req.rid}, hidden_len={dspark_hidden_len}, " + f"rid={decode_req.req.rid}, hidden_len={pd_hidden_len}, " f"pool_size={dspark_pool.size}. Increase " - "SGLANG_DSPARK_PD_HIDDEN_RECV_POOL_TOKENS." + "SGLANG_PD_HIDDEN_RECV_POOL_TOKENS." ) logger.error(message) prepare_abort( @@ -1249,13 +1249,13 @@ def pop_preallocated( failed_reqs.append(decode_req) indices_to_remove.add(i) continue - allocated_hidden_indices = dspark_pool.alloc(dspark_hidden_window_rows) + allocated_hidden_indices = dspark_pool.alloc(pd_hidden_window_rows) if allocated_hidden_indices is None: if prefix_len > 0: self.tree_cache.dec_lock_ref(decode_req.req.last_node) now = time.monotonic() if ( - now - self._last_dspark_hidden_recv_credit_warning_time + now - self._last_pd_hidden_recv_credit_warning_time > 30 ): logger.warning( @@ -1263,41 +1263,41 @@ def pop_preallocated( "rid=%s window_rows=%d hidden_len=%d free_rows=%d pool_rows=%d " "prealloc_queue=%d transfer_queue=%d", decode_req.req.rid, - dspark_hidden_window_rows, - dspark_hidden_len, + pd_hidden_window_rows, + pd_hidden_len, dspark_pool.available_size(), dspark_pool.size, len(self.queue), len(self.transfer_queue.queue), ) - self._last_dspark_hidden_recv_credit_warning_time = now + self._last_pd_hidden_recv_credit_warning_time = now continue - dspark_hidden_dst_indices_by_pp = {} + pd_hidden_dst_indices_by_pp = {} for pp_rank, pp_slice in pp_slices.items(): if int(pp_slice.get("slice_len", 0)) <= 0: pp_slice["dst_indices"] = [] - dspark_hidden_dst_indices_by_pp[int(pp_rank)] = [] + pd_hidden_dst_indices_by_pp[int(pp_rank)] = [] continue pp_slice["dst_indices"] = [ int(x) for x in allocated_hidden_indices ] - dspark_hidden_dst_indices_by_pp[int(pp_rank)] = [ + pd_hidden_dst_indices_by_pp[int(pp_rank)] = [ int(x) for x in allocated_hidden_indices ] - dspark_hidden_pp_slices = pp_slices - hidden_end = int(dspark_hidden_start + dspark_hidden_len) - decode_req.dspark_hidden_state = ( + pd_hidden_pp_slices = pp_slices + hidden_end = int(pd_hidden_start + pd_hidden_len) + decode_req.pd_hidden_state = ( PDHiddenRequestState.streaming_state( - int(dspark_hidden_start), hidden_end + int(pd_hidden_start), hidden_end ) - if dspark_hidden_streaming + if pd_hidden_streaming else PDHiddenRequestState.full( - int(dspark_hidden_start), hidden_end + int(pd_hidden_start), hidden_end ) ) if pp_size == 1: - dspark_hidden_dst_indices = dspark_hidden_dst_indices_by_pp.get(0) + pd_hidden_dst_indices = pd_hidden_dst_indices_by_pp.get(0) dst_kv_indices = self._pre_alloc( decode_req.req, @@ -1305,10 +1305,10 @@ def pop_preallocated( prefix_len, total_prefix_len, ) - decode_req.dspark_hidden_dst_indices = dspark_hidden_dst_indices - decode_req.dspark_hidden_dst_indices_by_pp = dspark_hidden_dst_indices_by_pp - decode_req.dspark_hidden_pp_slices = dspark_hidden_pp_slices - decode_req.dspark_hidden_start = dspark_hidden_start + decode_req.pd_hidden_dst_indices = pd_hidden_dst_indices + decode_req.pd_hidden_dst_indices_by_pp = pd_hidden_dst_indices_by_pp + decode_req.pd_hidden_pp_slices = pd_hidden_pp_slices + decode_req.pd_hidden_start = pd_hidden_start decode_req.prefix_match = prefix_match if self.scheduler.enable_decode_hicache: self._start_hicache_prefetch(decode_req.req, prefix_match) @@ -1423,11 +1423,11 @@ def _c128_state_payload(): state_indices.append(_swa_ring_payload()) elif st == StateType.C128_STATE: state_indices.append(_c128_state_payload()) - elif st == StateType.DSPARK_HIDDEN: + elif st == StateType.PD_HIDDEN: first_slice_indices = None - if dspark_hidden_dst_indices_by_pp: + if pd_hidden_dst_indices_by_pp: first_slice_indices = next( - iter(dspark_hidden_dst_indices_by_pp.values()) + iter(pd_hidden_dst_indices_by_pp.values()) ) state_indices.append( None @@ -1442,7 +1442,7 @@ def _c128_state_payload(): state_indices = None spec_metadata = None - if dspark_hidden_dst_indices_by_pp is not None: + if pd_hidden_dst_indices_by_pp is not None: model_runner = self.scheduler.tp_worker.model_runner spec_aux_config = getattr(model_runner, "spec_aux_config", None) target_layer_ids = ( @@ -1451,14 +1451,14 @@ def _c128_state_payload(): or [] ) spec_metadata = { - "dspark_hidden": True, - "streaming_hidden": bool(decode_req.dspark_hidden_state.streaming), + "pd_hidden": True, + "streaming_hidden": bool(decode_req.pd_hidden_state.streaming), "streaming_window_rows": int( max( ( len(indices) for indices in ( - dspark_hidden_dst_indices_by_pp or {} + pd_hidden_dst_indices_by_pp or {} ).values() ), default=0, @@ -1467,11 +1467,11 @@ def _c128_state_payload(): "decode_radix_cache_enabled": bool( self.scheduler.server_args.disaggregation_decode_enable_radix_cache ), - "hidden_start": int(dspark_hidden_start), - "hidden_len": int(dspark_hidden_len), + "hidden_start": int(pd_hidden_start), + "hidden_len": int(pd_hidden_len), "dst_indices": ( - [int(x) for x in dspark_hidden_dst_indices] - if dspark_hidden_dst_indices is not None + [int(x) for x in pd_hidden_dst_indices] + if pd_hidden_dst_indices is not None else [] ), "pp_slices": { @@ -1482,7 +1482,7 @@ def _c128_state_payload(): ], } for pp_rank, pp_slice in ( - dspark_hidden_pp_slices or {} + pd_hidden_pp_slices or {} ).items() }, "hidden_size": int(self.metadata_buffers.pd_hidden_pool.hidden_size), @@ -1994,9 +1994,9 @@ def extend(self, decode_reqs: List[DecodeRequest]) -> None: ): self.staging_handler.register_decode_req(dr.req.bootstrap_room, dr) - def _release_dspark_hidden_rows(self, decode_req: DecodeRequest) -> None: + def _release_pd_hidden_rows(self, decode_req: DecodeRequest) -> None: wait_ack_completions = getattr( - self.kv_manager, "wait_dspark_hidden_ack_completions", None + self.kv_manager, "wait_pd_hidden_ack_completions", None ) if wait_ack_completions is not None and not wait_ack_completions( decode_req.req.bootstrap_room @@ -2009,12 +2009,12 @@ def _release_dspark_hidden_rows(self, decode_req: DecodeRequest) -> None: ) return pop_acked_chunks = getattr( - self.kv_manager, "pop_dspark_hidden_acked_chunks", None + self.kv_manager, "pop_pd_hidden_acked_chunks", None ) if pop_acked_chunks is not None: pop_acked_chunks(decode_req.req.bootstrap_room) - indices_by_pp = decode_req.dspark_hidden_dst_indices_by_pp - indices = decode_req.dspark_hidden_dst_indices + indices_by_pp = decode_req.pd_hidden_dst_indices_by_pp + indices = decode_req.pd_hidden_dst_indices pool = getattr(self.metadata_buffers, "pd_hidden_pool", None) if pool is not None: if indices_by_pp is not None: @@ -2027,20 +2027,20 @@ def _release_dspark_hidden_rows(self, decode_req: DecodeRequest) -> None: pool.free(pp_indices) elif indices is not None: pool.free(indices) - decode_req.dspark_hidden_dst_indices = None - decode_req.dspark_hidden_dst_indices_by_pp = None - decode_req.dspark_hidden_pp_slices = None - decode_req.dspark_hidden_state.reset() + decode_req.pd_hidden_dst_indices = None + decode_req.pd_hidden_dst_indices_by_pp = None + decode_req.pd_hidden_pp_slices = None + decode_req.pd_hidden_state.reset() - def _consume_dspark_hidden_acked_chunks(self, decode_req: DecodeRequest) -> None: + def _consume_pd_hidden_acked_chunks(self, decode_req: DecodeRequest) -> None: pop_acked_chunks = getattr( - self.kv_manager, "pop_dspark_hidden_acked_chunks", None + self.kv_manager, "pop_pd_hidden_acked_chunks", None ) if pop_acked_chunks is None: return for chunk in pop_acked_chunks(decode_req.req.bootstrap_room): if chunk.get("is_last_hidden_chunk"): - decode_req.dspark_hidden_state.mark_hidden_done() + decode_req.pd_hidden_state.mark_hidden_done() def _commit_transfer_to_req(self, decode_req: DecodeRequest): idx = decode_req.metadata_buffer_index @@ -2088,7 +2088,7 @@ def _commit_transfer_to_req(self, decode_req: DecodeRequest): ) decode_req.kv_receiver.clear() decode_req.kv_receiver = None - self._release_dspark_hidden_rows(decode_req) + self._release_pd_hidden_rows(decode_req) return elif actual_room != expected_room: # Real corruption detected (mismatch) @@ -2108,7 +2108,7 @@ def _commit_transfer_to_req(self, decode_req: DecodeRequest): ) decode_req.kv_receiver.clear() decode_req.kv_receiver = None - self._release_dspark_hidden_rows(decode_req) + self._release_pd_hidden_rows(decode_req) return self._commit_hicache_local_restore_to_req(decode_req) @@ -2221,11 +2221,11 @@ def _poll_with_staging(self) -> list: server_args=self.scheduler.server_args, ) - def _drain_dspark_hidden_ready_chunks(self, decode_req: DecodeRequest) -> None: - hidden_state = decode_req.dspark_hidden_state + def _drain_pd_hidden_ready_chunks(self, decode_req: DecodeRequest) -> None: + hidden_state = decode_req.pd_hidden_state if not hidden_state.streaming: return - pop_chunks = getattr(self.kv_manager, "pop_dspark_hidden_ready_chunks", None) + pop_chunks = getattr(self.kv_manager, "pop_pd_hidden_ready_chunks", None) if pop_chunks is None: raise RuntimeError( "PD streaming hidden backend is missing ready chunk API." @@ -2288,7 +2288,7 @@ def _drain_dspark_hidden_ready_chunks(self, decode_req: DecodeRequest) -> None: hidden_chunk.hidden_start, ) submit_ack = getattr( - self.kv_manager, "submit_dspark_hidden_chunk_ack", None + self.kv_manager, "submit_pd_hidden_chunk_ack", None ) if submit_ack is None: raise RuntimeError( @@ -2344,8 +2344,8 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req for i, (decode_req, poll) in enumerate(zip(self.queue, polls)): if rids_to_check is not None and decode_req.req.rid not in rids_to_check: continue - self._consume_dspark_hidden_acked_chunks(decode_req) - self._drain_dspark_hidden_ready_chunks(decode_req) + self._consume_pd_hidden_acked_chunks(decode_req) + self._drain_pd_hidden_ready_chunks(decode_req) hicache_restore_status = decode_req.hicache_restore_status if ( @@ -2382,7 +2382,7 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req self.scheduler.hisparse_coordinator.request_finished(decode_req.req) # release pre-allocated kv cache, but don't insert into the tree since it's failed release_kv_cache(decode_req.req, self.tree_cache, is_insert=False) - self._release_dspark_hidden_rows(decode_req) + self._release_pd_hidden_rows(decode_req) decode_req.kv_receiver.clear() decode_req.kv_receiver = None indices_to_remove.add(i) @@ -2395,7 +2395,7 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req and hicache_restore_status == HiCacheRestoreResult.PENDING ): continue - hidden_state = decode_req.dspark_hidden_state + hidden_state = decode_req.pd_hidden_state hidden_state.mark_kv_done() if not hidden_state.request_done(): continue @@ -2439,7 +2439,7 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req # instead of the stale value, avoiding a false-positive mismatch. self.metadata_buffers.bootstrap_room[idx] = 0 self.req_to_metadata_buffer_idx_allocator.free(idx) - self._release_dspark_hidden_rows(self.queue[i]) + self._release_pd_hidden_rows(self.queue[i]) self.queue = [ entry for i, entry in enumerate(self.queue) if i not in indices_to_remove @@ -2450,7 +2450,7 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req def release_memory_occupation(self): """Clean up in-flight transfers before releasing GPU memory.""" for decode_req in self.queue: - self._release_dspark_hidden_rows(decode_req) + self._release_pd_hidden_rows(decode_req) self.queue.clear() def resume_memory_occupation(self): diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 864116d5817c..864e81ae3367 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -165,10 +165,10 @@ def from_zmq(cls, msg: List[bytes]): class MooncakeKVManager(CommonKVManager): AUX_DATA_HEADER = b"AUX_DATA" - DSPARK_HIDDEN_CHUNK_READY_HEADER = b"DSPARK_HIDDEN_CHUNK_READY" - DSPARK_HIDDEN_CHUNK_ACK_HEADER = b"DSPARK_HIDDEN_CHUNK_ACK" + PD_HIDDEN_CHUNK_READY_HEADER = b"PD_HIDDEN_CHUNK_READY" + PD_HIDDEN_CHUNK_ACK_HEADER = b"PD_HIDDEN_CHUNK_ACK" - def supports_dspark_hidden_streaming(self) -> bool: + def supports_pd_hidden_streaming(self) -> bool: return True def __init__( @@ -187,16 +187,16 @@ def __init__( self.session_failures = defaultdict(int) self.failed_sessions = set() self.session_lock = threading.Lock() - self.dspark_hidden_done_rooms = set() - self.dspark_hidden_done_lock = threading.Lock() - self.dspark_hidden_chunk_acks = defaultdict(int) - self.dspark_hidden_chunk_ack_cv = threading.Condition() - self.dspark_hidden_ack_waiters = {} - self.dspark_hidden_room_waiters = defaultdict(deque) - self.dspark_hidden_inflight_chunks = {} - self.dspark_hidden_inflight_lock = threading.Lock() - self.dspark_hidden_active_transfers = defaultdict(int) - self.dspark_hidden_active_cv = threading.Condition() + self.pd_hidden_done_rooms = set() + self.pd_hidden_done_lock = threading.Lock() + self.pd_hidden_chunk_acks = defaultdict(int) + self.pd_hidden_chunk_ack_cv = threading.Condition() + self.pd_hidden_ack_waiters = {} + self.pd_hidden_room_waiters = defaultdict(deque) + self.pd_hidden_inflight_chunks = {} + self.pd_hidden_inflight_lock = threading.Lock() + self.pd_hidden_active_transfers = defaultdict(int) + self.pd_hidden_active_cv = threading.Condition() # Determine the number of threads to use for kv sender cpu_count = os.cpu_count() transfer_thread_pool_size = ( @@ -256,19 +256,19 @@ def __init__( ).start() self.start_prefill_thread() elif self.disaggregation_mode == DisaggregationMode.DECODE: - self.dspark_hidden_ready_chunks: Dict[int, List[dict]] = defaultdict(list) - self.dspark_hidden_ready_lock = threading.Lock() - self.dspark_hidden_ack_completions = queue.SimpleQueue() - self.dspark_hidden_ack_pending_counts = defaultdict(int) - self.dspark_hidden_ack_completion_cv = threading.Condition() - self.dspark_hidden_acked_chunks: Dict[int, List[dict]] = defaultdict(list) - self.dspark_hidden_acked_lock = threading.Lock() - self.dspark_hidden_ack_wakeup_endpoint = ( + self.pd_hidden_ready_chunks: Dict[int, List[dict]] = defaultdict(list) + self.pd_hidden_ready_lock = threading.Lock() + self.pd_hidden_ack_completions = queue.SimpleQueue() + self.pd_hidden_ack_pending_counts = defaultdict(int) + self.pd_hidden_ack_completion_cv = threading.Condition() + self.pd_hidden_acked_chunks: Dict[int, List[dict]] = defaultdict(list) + self.pd_hidden_acked_lock = threading.Lock() + self.pd_hidden_ack_wakeup_endpoint = ( f"inproc://dspark-hidden-ack-{id(self)}" ) - self.dspark_hidden_ack_wakeup_receiver = self._zmq_ctx.socket(zmq.PULL) - self.dspark_hidden_ack_wakeup_receiver.bind( - self.dspark_hidden_ack_wakeup_endpoint + self.pd_hidden_ack_wakeup_receiver = self._zmq_ctx.socket(zmq.PULL) + self.pd_hidden_ack_wakeup_receiver.bind( + self.pd_hidden_ack_wakeup_endpoint ) self._staging_ctx = DecodeStagingContext() if self.enable_staging else None if self.enable_staging: @@ -277,20 +277,20 @@ def __init__( self._chunk_writer_counts: dict = defaultdict(lambda: defaultdict(list)) self.start_decode_thread() - def mark_dspark_hidden_request_done( + def mark_pd_hidden_request_done( self, bootstrap_room: int, state_indices: Optional[List] = None, ) -> None: - if not hasattr(self, "dspark_hidden_done_rooms"): + if not hasattr(self, "pd_hidden_done_rooms"): return - with self.dspark_hidden_done_lock: + with self.pd_hidden_done_lock: room = int(bootstrap_room) - if room in self.dspark_hidden_done_rooms: + if room in self.pd_hidden_done_rooms: return - self.dspark_hidden_done_rooms.add(room) + self.pd_hidden_done_rooms.add(room) pool = getattr(self, "pd_hidden_pool", None) - state_idx = self._dspark_hidden_state_index() + state_idx = self._pd_hidden_state_index() if ( pool is not None and state_indices is not None @@ -301,27 +301,27 @@ def mark_dspark_hidden_request_done( if indices is not None and len(indices) > 0: pool.free([int(idx) for idx in indices]) - def pop_dspark_hidden_request_done(self, bootstrap_room: int) -> bool: - if not hasattr(self, "dspark_hidden_done_rooms"): + def pop_pd_hidden_request_done(self, bootstrap_room: int) -> bool: + if not hasattr(self, "pd_hidden_done_rooms"): return False - with self.dspark_hidden_done_lock: + with self.pd_hidden_done_lock: room = int(bootstrap_room) - if room not in self.dspark_hidden_done_rooms: + if room not in self.pd_hidden_done_rooms: return False - self.dspark_hidden_done_rooms.remove(room) + self.pd_hidden_done_rooms.remove(room) return True - def mark_dspark_hidden_done( + def mark_pd_hidden_done( self, bootstrap_room: int, state_indices: Optional[List] = None, ) -> None: - self.mark_dspark_hidden_request_done(bootstrap_room, state_indices) + self.mark_pd_hidden_request_done(bootstrap_room, state_indices) - def pop_dspark_hidden_done(self, bootstrap_room: int) -> bool: - return self.pop_dspark_hidden_request_done(bootstrap_room) + def pop_pd_hidden_done(self, bootstrap_room: int) -> bool: + return self.pop_pd_hidden_request_done(bootstrap_room) - def notify_dspark_hidden_chunk_ready( + def notify_pd_hidden_chunk_ready( self, *, remote: str, @@ -336,7 +336,7 @@ def notify_dspark_hidden_chunk_ready( na = NetworkAddress(remote, dst_port) self._connect(na.to_tcp(), is_ipv6=na.is_ipv6).send_multipart( [ - self.DSPARK_HIDDEN_CHUNK_READY_HEADER, + self.PD_HIDDEN_CHUNK_READY_HEADER, str(room).encode("ascii"), str(prefill_rank).encode("ascii"), str(hidden_start).encode("ascii"), @@ -350,7 +350,7 @@ def notify_dspark_hidden_chunk_ready( ] ) - def wait_dspark_hidden_chunk_ack( + def wait_pd_hidden_chunk_ack( self, *, room: int, @@ -361,8 +361,8 @@ def wait_dspark_hidden_chunk_ack( ) -> bool: key = (int(room), int(prefill_rank), int(hidden_start)) deadline = time.monotonic() + float(timeout_s) - with self.dspark_hidden_chunk_ack_cv: - while self.dspark_hidden_chunk_acks.get(key, 0) < int(expected_count): + with self.pd_hidden_chunk_ack_cv: + while self.pd_hidden_chunk_acks.get(key, 0) < int(expected_count): if self.request_status.get(int(room)) == KVPoll.Failed: return False remaining = deadline - time.monotonic() @@ -374,13 +374,13 @@ def wait_dspark_hidden_chunk_ack( ) self.update_status(int(room), KVPoll.Failed) return False - self.dspark_hidden_chunk_ack_cv.wait(timeout=min(remaining, 1.0)) - self.dspark_hidden_chunk_acks[key] -= int(expected_count) - if self.dspark_hidden_chunk_acks[key] <= 0: - self.dspark_hidden_chunk_acks.pop(key, None) + self.pd_hidden_chunk_ack_cv.wait(timeout=min(remaining, 1.0)) + self.pd_hidden_chunk_acks[key] -= int(expected_count) + if self.pd_hidden_chunk_acks[key] <= 0: + self.pd_hidden_chunk_acks.pop(key, None) return True - def consume_dspark_hidden_chunk_ack( + def consume_pd_hidden_chunk_ack( self, *, room: int, @@ -389,15 +389,15 @@ def consume_dspark_hidden_chunk_ack( expected_count: int = 1, ) -> bool: key = (int(room), int(prefill_rank), int(hidden_start)) - with self.dspark_hidden_chunk_ack_cv: - if self.dspark_hidden_chunk_acks.get(key, 0) < int(expected_count): + with self.pd_hidden_chunk_ack_cv: + if self.pd_hidden_chunk_acks.get(key, 0) < int(expected_count): return False - self.dspark_hidden_chunk_acks[key] -= int(expected_count) - if self.dspark_hidden_chunk_acks[key] <= 0: - self.dspark_hidden_chunk_acks.pop(key, None) + self.pd_hidden_chunk_acks[key] -= int(expected_count) + if self.pd_hidden_chunk_acks[key] <= 0: + self.pd_hidden_chunk_acks.pop(key, None) return True - def park_dspark_hidden_chunk_for_ack( + def park_pd_hidden_chunk_for_ack( self, *, transfer_queue: FastQueue, @@ -411,39 +411,39 @@ def park_dspark_hidden_chunk_for_ack( Returns True when the chunk was parked. False means ACKs were already available and the caller can finish the chunk immediately. """ - if kv_chunk.dspark_hidden_start is None: + if kv_chunk.pd_hidden_start is None: return False key = ( int(kv_chunk.room), int(prefill_rank), - int(kv_chunk.dspark_hidden_start), + int(kv_chunk.pd_hidden_start), ) expected_count = int(expected_count) if expected_count <= 0: - kv_chunk.dspark_hidden_ack_ready = True + kv_chunk.pd_hidden_ack_ready = True return False - with self.dspark_hidden_chunk_ack_cv: - if self.dspark_hidden_chunk_acks.get(key, 0) >= expected_count: - self.dspark_hidden_chunk_acks[key] -= expected_count - if self.dspark_hidden_chunk_acks[key] <= 0: - self.dspark_hidden_chunk_acks.pop(key, None) - kv_chunk.dspark_hidden_ack_ready = True + with self.pd_hidden_chunk_ack_cv: + if self.pd_hidden_chunk_acks.get(key, 0) >= expected_count: + self.pd_hidden_chunk_acks[key] -= expected_count + if self.pd_hidden_chunk_acks[key] <= 0: + self.pd_hidden_chunk_acks.pop(key, None) + kv_chunk.pd_hidden_ack_ready = True return False - if key in self.dspark_hidden_ack_waiters: + if key in self.pd_hidden_ack_waiters: raise RuntimeError( "PD hidden ACK waiter already exists: " f"room={key[0]}, prefill_rank={key[1]}, hidden_start={key[2]}" ) - kv_chunk.dspark_hidden_ack_expected_count = expected_count - self.dspark_hidden_ack_waiters[key] = (transfer_queue, kv_chunk) + kv_chunk.pd_hidden_ack_expected_count = expected_count + self.pd_hidden_ack_waiters[key] = (transfer_queue, kv_chunk) def on_timeout() -> None: - with self.dspark_hidden_chunk_ack_cv: - waiter = self.dspark_hidden_ack_waiters.pop(key, None) + with self.pd_hidden_chunk_ack_cv: + waiter = self.pd_hidden_ack_waiters.pop(key, None) if waiter is None: return _, timed_out_chunk = waiter - timed_out_chunk.dspark_hidden_ack_timed_out = True + timed_out_chunk.pd_hidden_ack_timed_out = True self.record_failure( key[0], "Timed out waiting for PD hidden chunk ACK: " @@ -451,75 +451,75 @@ def on_timeout() -> None: ) self.update_status(key[0], KVPoll.Failed) transfer_queue.put(timed_out_chunk) - self._wake_dspark_hidden_ack_waiters(key[0]) + self._wake_pd_hidden_ack_waiters(key[0]) timer = threading.Timer(float(timeout_s), on_timeout) timer.daemon = True timer.start() return True - def _wake_dspark_hidden_ack_waiters(self, room: int) -> None: + def _wake_pd_hidden_ack_waiters(self, room: int) -> None: room = int(room) - with self.dspark_hidden_chunk_ack_cv: + with self.pd_hidden_chunk_ack_cv: waiters = [ - self.dspark_hidden_ack_waiters.pop(key) - for key in list(self.dspark_hidden_ack_waiters) + self.pd_hidden_ack_waiters.pop(key) + for key in list(self.pd_hidden_ack_waiters) if key[0] == room ] - room_waiters = list(self.dspark_hidden_room_waiters.pop(room, [])) + room_waiters = list(self.pd_hidden_room_waiters.pop(room, [])) for transfer_queue, kv_chunk in waiters: transfer_queue.put(kv_chunk) for transfer_queue, kv_chunk in room_waiters: transfer_queue.put(kv_chunk) - def _park_dspark_hidden_chunk_behind_room( + def _park_pd_hidden_chunk_behind_room( self, transfer_queue: FastQueue, kv_chunk: TransferKVChunk ) -> None: - with self.dspark_hidden_chunk_ack_cv: - self.dspark_hidden_room_waiters[int(kv_chunk.room)].append( + with self.pd_hidden_chunk_ack_cv: + self.pd_hidden_room_waiters[int(kv_chunk.room)].append( (transfer_queue, kv_chunk) ) - def _wake_next_dspark_hidden_room_waiter(self, room: int) -> None: - with self.dspark_hidden_chunk_ack_cv: - room_waiters = self.dspark_hidden_room_waiters.get(int(room)) + def _wake_next_pd_hidden_room_waiter(self, room: int) -> None: + with self.pd_hidden_chunk_ack_cv: + room_waiters = self.pd_hidden_room_waiters.get(int(room)) if not room_waiters: return transfer_queue, kv_chunk = room_waiters.popleft() if not room_waiters: - self.dspark_hidden_room_waiters.pop(int(room), None) + self.pd_hidden_room_waiters.pop(int(room), None) transfer_queue.put(kv_chunk) - def _handle_dspark_hidden_chunk_ack( + def _handle_pd_hidden_chunk_ack( self, room: int, prefill_rank: int, hidden_start: int ) -> None: key = (int(room), int(prefill_rank), int(hidden_start)) waiter_to_wake = None - with self.dspark_hidden_chunk_ack_cv: - self.dspark_hidden_chunk_acks[key] += 1 - waiter = self.dspark_hidden_ack_waiters.get(key) + with self.pd_hidden_chunk_ack_cv: + self.pd_hidden_chunk_acks[key] += 1 + waiter = self.pd_hidden_ack_waiters.get(key) if waiter is not None: _, kv_chunk = waiter - expected_count = kv_chunk.dspark_hidden_ack_expected_count - if self.dspark_hidden_chunk_acks[key] >= expected_count: - self.dspark_hidden_chunk_acks[key] -= expected_count - if self.dspark_hidden_chunk_acks[key] <= 0: - self.dspark_hidden_chunk_acks.pop(key, None) - self.dspark_hidden_ack_waiters.pop(key, None) - kv_chunk.dspark_hidden_ack_ready = True + expected_count = kv_chunk.pd_hidden_ack_expected_count + if self.pd_hidden_chunk_acks[key] >= expected_count: + self.pd_hidden_chunk_acks[key] -= expected_count + if self.pd_hidden_chunk_acks[key] <= 0: + self.pd_hidden_chunk_acks.pop(key, None) + self.pd_hidden_ack_waiters.pop(key, None) + kv_chunk.pd_hidden_ack_ready = True waiter_to_wake = waiter - self.dspark_hidden_chunk_ack_cv.notify_all() + self.pd_hidden_chunk_ack_cv.notify_all() if waiter_to_wake is not None: transfer_queue, kv_chunk = waiter_to_wake transfer_queue.put(kv_chunk) - def pop_dspark_hidden_ready_chunks(self, room: int) -> List[dict]: - if not hasattr(self, "dspark_hidden_ready_chunks"): + def pop_pd_hidden_ready_chunks(self, room: int) -> List[dict]: + if not hasattr(self, "pd_hidden_ready_chunks"): return [] - with self.dspark_hidden_ready_lock: - return self.dspark_hidden_ready_chunks.pop(int(room), []) + with self.pd_hidden_ready_lock: + return self.pd_hidden_ready_chunks.pop(int(room), []) - def ack_dspark_hidden_chunk( + def ack_pd_hidden_chunk( self, *, remote: str, @@ -531,14 +531,14 @@ def ack_dspark_hidden_chunk( na = NetworkAddress(remote, dst_port) self._connect(na.to_tcp(), is_ipv6=na.is_ipv6).send_multipart( [ - self.DSPARK_HIDDEN_CHUNK_ACK_HEADER, + self.PD_HIDDEN_CHUNK_ACK_HEADER, str(room).encode("ascii"), str(prefill_rank).encode("ascii"), str(hidden_start).encode("ascii"), ] ) - def submit_dspark_hidden_chunk_ack( + def submit_pd_hidden_chunk_ack( self, *, event, @@ -557,8 +557,8 @@ def submit_dspark_hidden_chunk_ack( "hidden_start": int(hidden_start), "is_last_hidden_chunk": bool(is_last_hidden_chunk), } - with self.dspark_hidden_ack_completion_cv: - self.dspark_hidden_ack_pending_counts[int(room)] += 1 + with self.pd_hidden_ack_completion_cv: + self.pd_hidden_ack_pending_counts[int(room)] += 1 def wait_for_injection() -> None: try: @@ -573,11 +573,11 @@ def wait_for_injection() -> None: hidden_start, ) completion["success"] = False - self.dspark_hidden_ack_completions.put(completion) + self.pd_hidden_ack_completions.put(completion) wakeup_sender = self._zmq_ctx.socket(zmq.PUSH) wakeup_sender.setsockopt(zmq.LINGER, 0) try: - wakeup_sender.connect(self.dspark_hidden_ack_wakeup_endpoint) + wakeup_sender.connect(self.pd_hidden_ack_wakeup_endpoint) wakeup_sender.send(b"ACK_READY") finally: wakeup_sender.close() @@ -588,25 +588,25 @@ def wait_for_injection() -> None: daemon=True, ).start() - def _drain_dspark_hidden_ack_completions(self) -> None: + def _drain_pd_hidden_ack_completions(self) -> None: while True: try: - completion = self.dspark_hidden_ack_completions.get_nowait() + completion = self.pd_hidden_ack_completions.get_nowait() except queue.Empty: return room = int(completion["room"]) try: if completion.pop("success"): - self.ack_dspark_hidden_chunk( + self.ack_pd_hidden_chunk( remote=completion["remote"], dst_port=int(completion["dst_port"]), room=room, prefill_rank=int(completion["prefill_rank"]), hidden_start=int(completion["hidden_start"]), ) - with self.dspark_hidden_acked_lock: - self.dspark_hidden_acked_chunks[room].append(completion) + with self.pd_hidden_acked_lock: + self.pd_hidden_acked_chunks[room].append(completion) else: self.record_failure( room, @@ -627,66 +627,66 @@ def _drain_dspark_hidden_ack_completions(self) -> None: ) self.update_status(room, KVPoll.Failed) finally: - with self.dspark_hidden_ack_completion_cv: - self.dspark_hidden_ack_pending_counts[room] -= 1 - if self.dspark_hidden_ack_pending_counts[room] <= 0: - self.dspark_hidden_ack_pending_counts.pop(room, None) - self.dspark_hidden_ack_completion_cv.notify_all() + with self.pd_hidden_ack_completion_cv: + self.pd_hidden_ack_pending_counts[room] -= 1 + if self.pd_hidden_ack_pending_counts[room] <= 0: + self.pd_hidden_ack_pending_counts.pop(room, None) + self.pd_hidden_ack_completion_cv.notify_all() - def pop_dspark_hidden_acked_chunks(self, room: int) -> List[dict]: - with self.dspark_hidden_acked_lock: - return self.dspark_hidden_acked_chunks.pop(int(room), []) + def pop_pd_hidden_acked_chunks(self, room: int) -> List[dict]: + with self.pd_hidden_acked_lock: + return self.pd_hidden_acked_chunks.pop(int(room), []) - def wait_dspark_hidden_ack_completions( + def wait_pd_hidden_ack_completions( self, room: int, timeout_s: float = 300.0 ) -> bool: room = int(room) deadline = time.monotonic() + float(timeout_s) - with self.dspark_hidden_ack_completion_cv: - while self.dspark_hidden_ack_pending_counts.get(room, 0) > 0: + with self.pd_hidden_ack_completion_cv: + while self.pd_hidden_ack_pending_counts.get(room, 0) > 0: remaining = deadline - time.monotonic() if remaining <= 0: return False - self.dspark_hidden_ack_completion_cv.wait(timeout=min(remaining, 1.0)) + self.pd_hidden_ack_completion_cv.wait(timeout=min(remaining, 1.0)) return True - def _begin_dspark_hidden_transfer(self, room: int) -> None: - if not hasattr(self, "dspark_hidden_active_cv"): + def _begin_pd_hidden_transfer(self, room: int) -> None: + if not hasattr(self, "pd_hidden_active_cv"): return - with self.dspark_hidden_active_cv: - self.dspark_hidden_active_transfers[int(room)] += 1 + with self.pd_hidden_active_cv: + self.pd_hidden_active_transfers[int(room)] += 1 - def _end_dspark_hidden_transfer(self, room: int) -> None: - if not hasattr(self, "dspark_hidden_active_cv"): + def _end_pd_hidden_transfer(self, room: int) -> None: + if not hasattr(self, "pd_hidden_active_cv"): return - with self.dspark_hidden_active_cv: + with self.pd_hidden_active_cv: room = int(room) - count = self.dspark_hidden_active_transfers.get(room, 0) - 1 + count = self.pd_hidden_active_transfers.get(room, 0) - 1 if count <= 0: - self.dspark_hidden_active_transfers.pop(room, None) + self.pd_hidden_active_transfers.pop(room, None) else: - self.dspark_hidden_active_transfers[room] = count - self.dspark_hidden_active_cv.notify_all() + self.pd_hidden_active_transfers[room] = count + self.pd_hidden_active_cv.notify_all() - def _wait_dspark_hidden_transfers_quiesced( + def _wait_pd_hidden_transfers_quiesced( self, room: int, timeout_s: float = 300.0 ) -> bool: - if not hasattr(self, "dspark_hidden_active_cv"): + if not hasattr(self, "pd_hidden_active_cv"): return True deadline = time.monotonic() + float(timeout_s) room = int(room) - with self.dspark_hidden_active_cv: - while self.dspark_hidden_active_transfers.get(room, 0) > 0: + with self.pd_hidden_active_cv: + while self.pd_hidden_active_transfers.get(room, 0) > 0: remaining = deadline - time.monotonic() if remaining <= 0: logger.error( "Timed out waiting for PD hidden transfers to quiesce: " "room=%s active=%s", room, - self.dspark_hidden_active_transfers.get(room, 0), + self.pd_hidden_active_transfers.get(room, 0), ) return False - self.dspark_hidden_active_cv.wait(timeout=min(remaining, 1.0)) + self.pd_hidden_active_cv.wait(timeout=min(remaining, 1.0)) return True def init_engine(self): @@ -1534,7 +1534,7 @@ def maybe_send_extra( src_indices = src_indices[: len(dst_indices_local)] else: dst_indices_local = dst_indices_local[: len(src_indices)] - if st == StateType.DSPARK_HIDDEN and dynamic_dst: + if st == StateType.PD_HIDDEN and dynamic_dst: row_chunks = dynamic_dst.get("row_chunks") or [ {"row_start": 0, "row_len": len(src_indices)} ] @@ -1624,45 +1624,45 @@ def maybe_send_extra( ) return rc - def _dspark_hidden_state_index(self) -> Optional[int]: + def _pd_hidden_state_index(self) -> Optional[int]: for idx, state_type in enumerate(getattr(self.kv_args, "state_types", [])): - if state_type == StateType.DSPARK_HIDDEN: + if state_type == StateType.PD_HIDDEN: return idx return None - def _has_dspark_hidden_state(self, state_indices: Optional[List]) -> bool: - idx = self._dspark_hidden_state_index() + def _has_pd_hidden_state(self, state_indices: Optional[List]) -> bool: + idx = self._pd_hidden_state_index() if idx is None or not state_indices or idx >= len(state_indices): return False indices = state_indices[idx] return indices is not None and len(indices) > 0 - def _without_dspark_hidden_state( + def _without_pd_hidden_state( self, state_indices: Optional[List] ) -> Optional[List]: - idx = self._dspark_hidden_state_index() + idx = self._pd_hidden_state_index() if idx is None or not state_indices or idx >= len(state_indices): return state_indices ret = list(state_indices) ret[idx] = None return ret - def _dspark_hidden_release_state_indices( + def _pd_hidden_release_state_indices( self, kv_chunk: TransferKVChunk ) -> Optional[List]: - release_indices = kv_chunk.dspark_hidden_release_indices + release_indices = kv_chunk.pd_hidden_release_indices if not release_indices: return kv_chunk.state_indices - idx = self._dspark_hidden_state_index() + idx = self._pd_hidden_state_index() if idx is None or not kv_chunk.state_indices or idx >= len(kv_chunk.state_indices): return kv_chunk.state_indices ret = list(kv_chunk.state_indices) ret[idx] = [int(x) for x in release_indices] return ret - def _free_dspark_hidden_state_indices(self, state_indices: Optional[List]) -> None: + def _free_pd_hidden_state_indices(self, state_indices: Optional[List]) -> None: pool = getattr(self, "pd_hidden_pool", None) - state_idx = self._dspark_hidden_state_index() + state_idx = self._pd_hidden_state_index() if ( pool is None or state_idx is None @@ -1674,19 +1674,19 @@ def _free_dspark_hidden_state_indices(self, state_indices: Optional[List]) -> No if indices is not None and len(indices) > 0: pool.free([int(idx) for idx in indices]) - def _free_dspark_hidden_chunk_rows(self, kv_chunk: TransferKVChunk) -> None: - self._free_dspark_hidden_state_indices( - self._dspark_hidden_release_state_indices(kv_chunk) + def _free_pd_hidden_chunk_rows(self, kv_chunk: TransferKVChunk) -> None: + self._free_pd_hidden_state_indices( + self._pd_hidden_release_state_indices(kv_chunk) ) - def _send_dspark_hidden_packet( + def _send_pd_hidden_packet( self, req: TransferInfo, prefill_state_indices: List, packet_idx: int, executor: concurrent.futures.ThreadPoolExecutor, ) -> Tuple[int, bool]: - state_idx = self._dspark_hidden_state_index() + state_idx = self._pd_hidden_state_index() if state_idx is None or state_idx >= len(prefill_state_indices): return 0, True indices = prefill_state_indices[state_idx] @@ -1706,7 +1706,7 @@ def _send_dspark_hidden_packet( dst_indices = dst_indices[: len(src_indices)] if len(src_indices) != len(dst_indices): raise RuntimeError( - "DSPARK_HIDDEN state index length mismatch: " + "PD_HIDDEN state index length mismatch: " f"room={req.room}, prefill={len(src_indices)}, " f"dst={len(dst_indices)}" ) @@ -1726,7 +1726,7 @@ def _send_dspark_hidden_packet( prefill_data_indices=src_indices, dst_data_indices=dst_indices, executor=executor, - state_type=StateType.DSPARK_HIDDEN, + state_type=StateType.PD_HIDDEN, ) return rc, True row_chunks = dynamic_dst.get("row_chunks") or [] @@ -1758,7 +1758,7 @@ def _send_dspark_hidden_packet( prefill_data_indices=src_indices[row_start:row_end], dst_data_indices=np.arange(row_len, dtype=np.int32), executor=executor, - state_type=StateType.DSPARK_HIDDEN, + state_type=StateType.PD_HIDDEN, ) return rc, packet_idx + 1 >= len(row_chunks) @@ -1922,25 +1922,25 @@ def transfer_worker( current_status = self.request_status.get(kv_chunk.room) if current_status is None or current_status == KVPoll.Failed: if current_status == KVPoll.Failed: - self._wake_dspark_hidden_ack_waiters(kv_chunk.room) + self._wake_pd_hidden_ack_waiters(kv_chunk.room) logger.debug( f"Skipping chunk for room {kv_chunk.room} because it has already failed or been aborted" ) - if kv_chunk.dspark_hidden_start is not None: - with self.dspark_hidden_inflight_lock: - self.dspark_hidden_inflight_chunks.pop( + if kv_chunk.pd_hidden_start is not None: + with self.pd_hidden_inflight_lock: + self.pd_hidden_inflight_chunks.pop( kv_chunk.room, None ) if ( - not kv_chunk.dspark_hidden_sent - and self._has_dspark_hidden_state(kv_chunk.state_indices) + not kv_chunk.pd_hidden_sent + and self._has_pd_hidden_state(kv_chunk.state_indices) ): - if kv_chunk.dspark_hidden_start is not None: - self._free_dspark_hidden_chunk_rows(kv_chunk) + if kv_chunk.pd_hidden_start is not None: + self._free_pd_hidden_chunk_rows(kv_chunk) else: - self.mark_dspark_hidden_request_done( + self.mark_pd_hidden_request_done( kv_chunk.room, - self._dspark_hidden_release_state_indices(kv_chunk), + self._pd_hidden_release_state_indices(kv_chunk), ) if self.enable_trace: kv_chunk.trace_ctx.trace_slice_end( @@ -1963,8 +1963,8 @@ def transfer_worker( ) polls = [] dst_ranks_infos = [] - dspark_hidden_expected = 0 - dspark_hidden_done_count = 0 + pd_hidden_expected = 0 + pd_hidden_done_count = 0 # Unique id per prefill sender so decode's response set size matches expected_response_num. prefill_unique_rank = ( self.attn_tp_rank * (self.pp_size * self.attn_cp_size) @@ -1974,14 +1974,14 @@ def transfer_worker( # When staging transfer is not yet ready (watermark/allocation pending), # the chunk is re-enqueued and we break out of the req loop to retry later. staging_deferred = False - dspark_hidden_deferred = False - dspark_hidden_failed = False + pd_hidden_deferred = False + pd_hidden_failed = False if ( - not kv_chunk.dspark_hidden_sent + not kv_chunk.pd_hidden_sent and kv_chunk.state_indices - and self._has_dspark_hidden_state(kv_chunk.state_indices) + and self._has_pd_hidden_state(kv_chunk.state_indices) ): - state_idx = self._dspark_hidden_state_index() + state_idx = self._pd_hidden_state_index() for req in reqs_to_be_processed: if req.is_dummy: continue @@ -1992,52 +1992,52 @@ def transfer_worker( target_rank_registration_info ) if not skip_state: - dspark_hidden_expected += 1 + pd_hidden_expected += 1 waiting_for_ack = ( - kv_chunk.dspark_hidden_start is not None - and kv_chunk.dspark_hidden_ready_sent + kv_chunk.pd_hidden_start is not None + and kv_chunk.pd_hidden_ready_sent ) hidden_inflight_key = ( - (prefill_unique_rank, int(kv_chunk.dspark_hidden_start)) - if kv_chunk.dspark_hidden_start is not None + (prefill_unique_rank, int(kv_chunk.pd_hidden_start)) + if kv_chunk.pd_hidden_start is not None else None ) if ( hidden_inflight_key is not None - and not kv_chunk.dspark_hidden_ready_sent + and not kv_chunk.pd_hidden_ready_sent ): - with self.dspark_hidden_inflight_lock: - inflight_key = self.dspark_hidden_inflight_chunks.get( + with self.pd_hidden_inflight_lock: + inflight_key = self.pd_hidden_inflight_chunks.get( kv_chunk.room ) if inflight_key is not None and inflight_key != hidden_inflight_key: - self._park_dspark_hidden_chunk_behind_room( + self._park_pd_hidden_chunk_behind_room( queue, kv_chunk ) continue ack_ready = False if waiting_for_ack: - ack_ready = kv_chunk.dspark_hidden_ack_ready + ack_ready = kv_chunk.pd_hidden_ack_ready if ack_ready: - dspark_hidden_done_count = dspark_hidden_expected - self._free_dspark_hidden_chunk_rows(kv_chunk) + pd_hidden_done_count = pd_hidden_expected + self._free_pd_hidden_chunk_rows(kv_chunk) if hidden_inflight_key is not None: - with self.dspark_hidden_inflight_lock: + with self.pd_hidden_inflight_lock: if ( - self.dspark_hidden_inflight_chunks.get( + self.pd_hidden_inflight_chunks.get( kv_chunk.room ) == hidden_inflight_key ): - self.dspark_hidden_inflight_chunks.pop( + self.pd_hidden_inflight_chunks.pop( kv_chunk.room, None ) - self._wake_next_dspark_hidden_room_waiter( + self._wake_next_pd_hidden_room_waiter( kv_chunk.room ) - elif kv_chunk.dspark_hidden_ack_timed_out: - dspark_hidden_failed = True + elif kv_chunk.pd_hidden_ack_timed_out: + pd_hidden_failed = True if not ack_ready and not waiting_for_ack: for req in reqs_to_be_processed: @@ -2057,18 +2057,18 @@ def transfer_worker( ) if skip_state: continue - self._begin_dspark_hidden_transfer(kv_chunk.room) + self._begin_pd_hidden_transfer(kv_chunk.room) try: - ret, dspark_hidden_done = ( - self._send_dspark_hidden_packet( + ret, pd_hidden_done = ( + self._send_pd_hidden_packet( req, kv_chunk.state_indices, - kv_chunk.dspark_hidden_packet_idx, + kv_chunk.pd_hidden_packet_idx, executor, ) ) finally: - self._end_dspark_hidden_transfer(kv_chunk.room) + self._end_pd_hidden_transfer(kv_chunk.room) if ret != 0: with self.session_lock: self.session_failures[ @@ -2087,15 +2087,15 @@ def transfer_worker( self.record_failure( kv_chunk.room, "Failed to send PD hidden packet " - f"{kv_chunk.dspark_hidden_packet_idx} of " + f"{kv_chunk.pd_hidden_packet_idx} of " f"{kv_chunk.room} to " f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}", ) self.update_status(kv_chunk.room, KVPoll.Failed) - self._wake_dspark_hidden_ack_waiters(kv_chunk.room) - if kv_chunk.dspark_hidden_start is not None: - with self.dspark_hidden_inflight_lock: - self.dspark_hidden_inflight_chunks.pop( + self._wake_pd_hidden_ack_waiters(kv_chunk.room) + if kv_chunk.pd_hidden_start is not None: + with self.pd_hidden_inflight_lock: + self.pd_hidden_inflight_chunks.pop( kv_chunk.room, None ) self.sync_status_to_decode_endpoint( @@ -2105,22 +2105,22 @@ def transfer_worker( KVPoll.Failed, prefill_unique_rank, ) - if kv_chunk.dspark_hidden_start is not None: - self._free_dspark_hidden_chunk_rows(kv_chunk) + if kv_chunk.pd_hidden_start is not None: + self._free_pd_hidden_chunk_rows(kv_chunk) else: - self.mark_dspark_hidden_request_done( + self.mark_pd_hidden_request_done( kv_chunk.room, - self._dspark_hidden_release_state_indices( + self._pd_hidden_release_state_indices( kv_chunk ), ) - dspark_hidden_failed = True + pd_hidden_failed = True break - if not dspark_hidden_done: - dspark_hidden_deferred = True + if not pd_hidden_done: + pd_hidden_deferred = True continue if ( - kv_chunk.dspark_hidden_start is not None + kv_chunk.pd_hidden_start is not None and state_idx is not None and state_idx < len(kv_chunk.state_indices) ): @@ -2129,26 +2129,26 @@ def transfer_worker( if state_idx < len(req.dst_state_indices) else [] ) - row_len = int(kv_chunk.dspark_hidden_row_len) + row_len = int(kv_chunk.pd_hidden_row_len) chunk_dst_indices = [ int(x) for x in dst_indices[:row_len] ] - self.notify_dspark_hidden_chunk_ready( + self.notify_pd_hidden_chunk_ready( remote=req.endpoint, dst_port=req.dst_port, room=req.room, prefill_rank=prefill_unique_rank, - hidden_start=int(kv_chunk.dspark_hidden_start), + hidden_start=int(kv_chunk.pd_hidden_start), row_len=row_len, is_last_hidden_chunk=bool( - kv_chunk.dspark_hidden_is_last_chunk + kv_chunk.pd_hidden_is_last_chunk ), dst_indices=chunk_dst_indices, ) continue - dspark_hidden_done_count += 1 + pd_hidden_done_count += 1 - if dspark_hidden_failed: + if pd_hidden_failed: continue if waiting_for_ack and not ack_ready: # A parked chunk is only re-enqueued by ACK/abort/timeout. @@ -2156,54 +2156,54 @@ def transfer_worker( # duplicate wakeup; avoid busy requeueing it. continue if ( - kv_chunk.dspark_hidden_start is not None - and not kv_chunk.dspark_hidden_ready_sent + kv_chunk.pd_hidden_start is not None + and not kv_chunk.pd_hidden_ready_sent ): if hidden_inflight_key is not None: - with self.dspark_hidden_inflight_lock: - self.dspark_hidden_inflight_chunks[ + with self.pd_hidden_inflight_lock: + self.pd_hidden_inflight_chunks[ kv_chunk.room ] = hidden_inflight_key - kv_chunk.dspark_hidden_ready_sent = True - if self.park_dspark_hidden_chunk_for_ack( + kv_chunk.pd_hidden_ready_sent = True + if self.park_pd_hidden_chunk_for_ack( transfer_queue=queue, kv_chunk=kv_chunk, prefill_rank=prefill_unique_rank, - expected_count=dspark_hidden_expected, + expected_count=pd_hidden_expected, ): continue ack_ready = True - dspark_hidden_done_count = dspark_hidden_expected - self._free_dspark_hidden_chunk_rows(kv_chunk) + pd_hidden_done_count = pd_hidden_expected + self._free_pd_hidden_chunk_rows(kv_chunk) if hidden_inflight_key is not None: - with self.dspark_hidden_inflight_lock: - self.dspark_hidden_inflight_chunks.pop( + with self.pd_hidden_inflight_lock: + self.pd_hidden_inflight_chunks.pop( kv_chunk.room, None ) - self._wake_next_dspark_hidden_room_waiter( + self._wake_next_pd_hidden_room_waiter( kv_chunk.room ) current_status = self.request_status.get(kv_chunk.room) if ( - dspark_hidden_expected > 0 - and dspark_hidden_done_count == dspark_hidden_expected + pd_hidden_expected > 0 + and pd_hidden_done_count == pd_hidden_expected and current_status is not None and current_status != KVPoll.Failed ): - kv_chunk.dspark_hidden_sent = True - if kv_chunk.dspark_hidden_start is not None: - if kv_chunk.dspark_hidden_is_last_chunk: - self.mark_dspark_hidden_request_done( + kv_chunk.pd_hidden_sent = True + if kv_chunk.pd_hidden_start is not None: + if kv_chunk.pd_hidden_is_last_chunk: + self.mark_pd_hidden_request_done( kv_chunk.room, None, ) else: - self.mark_dspark_hidden_request_done( + self.mark_pd_hidden_request_done( kv_chunk.room, - self._dspark_hidden_release_state_indices(kv_chunk), + self._pd_hidden_release_state_indices(kv_chunk), ) - if dspark_hidden_deferred: - kv_chunk.dspark_hidden_packet_idx += 1 + if pd_hidden_deferred: + kv_chunk.pd_hidden_packet_idx += 1 queue.put(kv_chunk) continue @@ -2227,17 +2227,17 @@ def transfer_worker( ) if ( kv_chunk.is_last_chunk - and not kv_chunk.dspark_hidden_sent - and self._has_dspark_hidden_state( + and not kv_chunk.pd_hidden_sent + and self._has_pd_hidden_state( kv_chunk.state_indices ) ): - if kv_chunk.dspark_hidden_start is not None: - self._free_dspark_hidden_chunk_rows(kv_chunk) + if kv_chunk.pd_hidden_start is not None: + self._free_pd_hidden_chunk_rows(kv_chunk) else: - self.mark_dspark_hidden_request_done( + self.mark_pd_hidden_request_done( kv_chunk.room, - self._dspark_hidden_release_state_indices( + self._pd_hidden_release_state_indices( kv_chunk ), ) @@ -2326,7 +2326,7 @@ def transfer_worker( f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}", ) self.update_status(kv_chunk.room, KVPoll.Failed) - self._wake_dspark_hidden_ack_waiters(kv_chunk.room) + self._wake_pd_hidden_ack_waiters(kv_chunk.room) self.sync_status_to_decode_endpoint( req.endpoint, req.dst_port, @@ -2337,26 +2337,26 @@ def transfer_worker( break if kv_chunk.is_last_chunk: - has_dspark_hidden = ( + has_pd_hidden = ( kv_chunk.state_indices - and not kv_chunk.dspark_hidden_sent + and not kv_chunk.pd_hidden_sent and not skip_state - and self._has_dspark_hidden_state(kv_chunk.state_indices) + and self._has_pd_hidden_state(kv_chunk.state_indices) ) - if has_dspark_hidden: - dspark_hidden_expected += 1 - self._begin_dspark_hidden_transfer(kv_chunk.room) + if has_pd_hidden: + pd_hidden_expected += 1 + self._begin_pd_hidden_transfer(kv_chunk.room) try: - ret, dspark_hidden_done = ( - self._send_dspark_hidden_packet( + ret, pd_hidden_done = ( + self._send_pd_hidden_packet( req, kv_chunk.state_indices, - kv_chunk.dspark_hidden_packet_idx, + kv_chunk.pd_hidden_packet_idx, executor, ) ) finally: - self._end_dspark_hidden_transfer(kv_chunk.room) + self._end_pd_hidden_transfer(kv_chunk.room) if ret != 0: with self.session_lock: self.session_failures[ @@ -2377,12 +2377,12 @@ def transfer_worker( self.record_failure( kv_chunk.room, "Failed to send PD hidden packet " - f"{kv_chunk.dspark_hidden_packet_idx} of " + f"{kv_chunk.pd_hidden_packet_idx} of " f"{kv_chunk.room} to " f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}", ) self.update_status(kv_chunk.room, KVPoll.Failed) - self._wake_dspark_hidden_ack_waiters( + self._wake_pd_hidden_ack_waiters( kv_chunk.room ) self.sync_status_to_decode_endpoint( @@ -2392,25 +2392,25 @@ def transfer_worker( KVPoll.Failed, prefill_unique_rank, ) - if kv_chunk.dspark_hidden_start is not None: - self._free_dspark_hidden_chunk_rows(kv_chunk) + if kv_chunk.pd_hidden_start is not None: + self._free_pd_hidden_chunk_rows(kv_chunk) else: - self.mark_dspark_hidden_request_done( + self.mark_pd_hidden_request_done( kv_chunk.room, - self._dspark_hidden_release_state_indices( + self._pd_hidden_release_state_indices( kv_chunk ), ) break - if not dspark_hidden_done: - dspark_hidden_deferred = True + if not pd_hidden_done: + pd_hidden_deferred = True continue - dspark_hidden_done_count += 1 + pd_hidden_done_count += 1 if kv_chunk.state_indices and not skip_state: self.maybe_send_extra( req, - self._without_dspark_hidden_state( + self._without_pd_hidden_state( kv_chunk.state_indices ), executor, @@ -2466,25 +2466,25 @@ def transfer_worker( kv_chunk.kv_sent = True if ( kv_chunk.is_last_chunk - and not kv_chunk.dspark_hidden_sent - and dspark_hidden_expected > 0 - and dspark_hidden_done_count == dspark_hidden_expected + and not kv_chunk.pd_hidden_sent + and pd_hidden_expected > 0 + and pd_hidden_done_count == pd_hidden_expected and current_status is not None and current_status != KVPoll.Failed ): - if kv_chunk.dspark_hidden_start is not None: - if kv_chunk.dspark_hidden_is_last_chunk: - self.mark_dspark_hidden_request_done( + if kv_chunk.pd_hidden_start is not None: + if kv_chunk.pd_hidden_is_last_chunk: + self.mark_pd_hidden_request_done( kv_chunk.room, - self._dspark_hidden_release_state_indices(kv_chunk), + self._pd_hidden_release_state_indices(kv_chunk), ) else: - self.mark_dspark_hidden_request_done( + self.mark_pd_hidden_request_done( kv_chunk.room, - self._dspark_hidden_release_state_indices(kv_chunk), + self._pd_hidden_release_state_indices(kv_chunk), ) - if dspark_hidden_deferred: - kv_chunk.dspark_hidden_packet_idx += 1 + if pd_hidden_deferred: + kv_chunk.pd_hidden_packet_idx += 1 queue.put(kv_chunk) continue @@ -2506,11 +2506,11 @@ def bootstrap_thread(): # KVPoll.Bootstrapping -> KVPoll.WaitingForInput while True: waiting_req_bytes = self.server_socket.recv_multipart() - if waiting_req_bytes[0] == MooncakeKVManager.DSPARK_HIDDEN_CHUNK_ACK_HEADER: + if waiting_req_bytes[0] == MooncakeKVManager.PD_HIDDEN_CHUNK_ACK_HEADER: room = int(waiting_req_bytes[1].decode("ascii")) prefill_rank = int(waiting_req_bytes[2].decode("ascii")) hidden_start = int(waiting_req_bytes[3].decode("ascii")) - self._handle_dspark_hidden_chunk_ack( + self._handle_pd_hidden_chunk_ack( room, prefill_rank, hidden_start ) continue @@ -2542,7 +2542,7 @@ def bootstrap_thread(): and self.check_status(room_to_be_aborted) != KVPoll.Success ): self.update_status(room_to_be_aborted, KVPoll.Failed) - self._wake_dspark_hidden_ack_waiters(room_to_be_aborted) + self._wake_pd_hidden_ack_waiters(room_to_be_aborted) logger.debug( f"Received abort notification for room {room_to_be_aborted}, " f"marked as Failed" @@ -2552,7 +2552,7 @@ def bootstrap_thread(): f"Received abort notification for room {room_to_be_aborted}, " f"ignoring (already completed or unknown)" ) - self._wait_dspark_hidden_transfers_quiesced(room_to_be_aborted) + self._wait_pd_hidden_transfers_quiesced(room_to_be_aborted) # Send ACK back to decode endpoint try: na = NetworkAddress(decode_ip, decode_port) @@ -2617,12 +2617,12 @@ def bootstrap_thread(): info.spec_metadata for info in self.transfer_infos[room].values() if info.spec_metadata - and info.spec_metadata.get("dspark_hidden") + and info.spec_metadata.get("pd_hidden") ), None, ) if dspark_meta: - self.req_to_dspark_hidden_meta[room] = dspark_meta + self.req_to_pd_hidden_meta[room] = dspark_meta self.update_status(room, KVPoll.WaitingForInput) threading.Thread(target=bootstrap_thread).start() @@ -2631,23 +2631,23 @@ def start_decode_thread(self): def decode_thread(): poller = zmq.Poller() poller.register(self.server_socket, zmq.POLLIN) - poller.register(self.dspark_hidden_ack_wakeup_receiver, zmq.POLLIN) + poller.register(self.pd_hidden_ack_wakeup_receiver, zmq.POLLIN) while True: events = dict(poller.poll()) - if self.dspark_hidden_ack_wakeup_receiver in events: + if self.pd_hidden_ack_wakeup_receiver in events: while True: try: - self.dspark_hidden_ack_wakeup_receiver.recv(zmq.NOBLOCK) + self.pd_hidden_ack_wakeup_receiver.recv(zmq.NOBLOCK) except zmq.Again: break - self._drain_dspark_hidden_ack_completions() + self._drain_pd_hidden_ack_completions() if self.server_socket not in events: continue msg = self.server_socket.recv_multipart() if msg[0] == MooncakeKVManager.AUX_DATA_HEADER: self._handle_aux_data(msg) continue - if msg[0] == MooncakeKVManager.DSPARK_HIDDEN_CHUNK_READY_HEADER: + if msg[0] == MooncakeKVManager.PD_HIDDEN_CHUNK_READY_HEADER: room = int(msg[1].decode("ascii")) prefill_rank = int(msg[2].decode("ascii")) hidden_start = int(msg[3].decode("ascii")) @@ -2660,8 +2660,8 @@ def decode_thread(): ) ack_host = msg[7].decode("ascii") ack_port = int(msg[8].decode("ascii")) - with self.dspark_hidden_ready_lock: - self.dspark_hidden_ready_chunks[room].append( + with self.pd_hidden_ready_lock: + self.pd_hidden_ready_chunks[room].append( { "room": room, "prefill_rank": prefill_rank, @@ -2674,14 +2674,14 @@ def decode_thread(): } ) continue - if msg[0] == MooncakeKVManager.DSPARK_HIDDEN_CHUNK_ACK_HEADER: + if msg[0] == MooncakeKVManager.PD_HIDDEN_CHUNK_ACK_HEADER: room = int(msg[1].decode("ascii")) prefill_rank = int(msg[2].decode("ascii")) hidden_start = int(msg[3].decode("ascii")) key = (room, prefill_rank, hidden_start) - with self.dspark_hidden_chunk_ack_cv: - self.dspark_hidden_chunk_acks[key] += 1 - self.dspark_hidden_chunk_ack_cv.notify_all() + with self.pd_hidden_chunk_ack_cv: + self.pd_hidden_chunk_acks[key] += 1 + self.pd_hidden_chunk_ack_cv.notify_all() continue # Staging: prefill notifies a chunk written to staging buffer @@ -2758,10 +2758,10 @@ def add_transfer_request( state_indices: Optional[List] = None, trace_ctx: Optional[Union[TraceReqContext, TraceNullContext]] = None, source_event=None, - dspark_hidden_start: Optional[int] = None, - dspark_hidden_row_len: int = 0, - dspark_hidden_is_last_chunk: bool = False, - dspark_hidden_release_indices: Optional[List[int]] = None, + pd_hidden_start: Optional[int] = None, + pd_hidden_row_len: int = 0, + pd_hidden_is_last_chunk: bool = False, + pd_hidden_release_indices: Optional[List[int]] = None, ): assert self.disaggregation_mode == DisaggregationMode.PREFILL assert not is_last_chunk or (is_last_chunk and aux_index is not None) @@ -2800,10 +2800,10 @@ def add_transfer_request( prefill_aux_index=aux_index, state_indices=state_indices, source_event=source_event, - dspark_hidden_start=dspark_hidden_start, - dspark_hidden_row_len=dspark_hidden_row_len, - dspark_hidden_is_last_chunk=dspark_hidden_is_last_chunk, - dspark_hidden_release_indices=dspark_hidden_release_indices, + pd_hidden_start=pd_hidden_start, + pd_hidden_row_len=pd_hidden_row_len, + pd_hidden_is_last_chunk=pd_hidden_is_last_chunk, + pd_hidden_release_indices=pd_hidden_release_indices, trace_ctx=trace_ctx, ) ) @@ -2878,20 +2878,20 @@ def __init__( self.conclude_state = None self.init_time = time.time() self._source_event = None - self._dspark_hidden_chunk_meta = None + self._pd_hidden_chunk_meta = None self._init_trace_ctx() def set_source_event(self, source_event) -> None: self._source_event = source_event - def set_dspark_hidden_chunk_meta( + def set_pd_hidden_chunk_meta( self, hidden_start: int, row_len: int, is_last_hidden_chunk: bool, release_indices: Optional[List[int]] = None, ) -> None: - self._dspark_hidden_chunk_meta = ( + self._pd_hidden_chunk_meta = ( int(hidden_start), int(row_len), bool(is_last_hidden_chunk), @@ -2911,8 +2911,8 @@ def send( self._source_event = None return - dspark_hidden_chunk_meta = self._dspark_hidden_chunk_meta - self._dspark_hidden_chunk_meta = None + pd_hidden_chunk_meta = self._pd_hidden_chunk_meta + self._pd_hidden_chunk_meta = None if not is_last_chunk: source_event = self._source_event self._source_event = None @@ -2924,17 +2924,17 @@ def send( state_indices=state_indices, source_event=source_event, trace_ctx=self.trace_ctx.copy_for_thread(), - dspark_hidden_start=( - dspark_hidden_chunk_meta[0] if dspark_hidden_chunk_meta else None + pd_hidden_start=( + pd_hidden_chunk_meta[0] if pd_hidden_chunk_meta else None ), - dspark_hidden_row_len=( - dspark_hidden_chunk_meta[1] if dspark_hidden_chunk_meta else 0 + pd_hidden_row_len=( + pd_hidden_chunk_meta[1] if pd_hidden_chunk_meta else 0 ), - dspark_hidden_is_last_chunk=( - dspark_hidden_chunk_meta[2] if dspark_hidden_chunk_meta else False + pd_hidden_is_last_chunk=( + pd_hidden_chunk_meta[2] if pd_hidden_chunk_meta else False ), - dspark_hidden_release_indices=( - dspark_hidden_chunk_meta[3] if dspark_hidden_chunk_meta else None + pd_hidden_release_indices=( + pd_hidden_chunk_meta[3] if pd_hidden_chunk_meta else None ), ) else: @@ -2949,17 +2949,17 @@ def send( state_indices=state_indices, source_event=source_event, trace_ctx=self.trace_ctx.copy_for_thread(), - dspark_hidden_start=( - dspark_hidden_chunk_meta[0] if dspark_hidden_chunk_meta else None + pd_hidden_start=( + pd_hidden_chunk_meta[0] if pd_hidden_chunk_meta else None ), - dspark_hidden_row_len=( - dspark_hidden_chunk_meta[1] if dspark_hidden_chunk_meta else 0 + pd_hidden_row_len=( + pd_hidden_chunk_meta[1] if pd_hidden_chunk_meta else 0 ), - dspark_hidden_is_last_chunk=( - dspark_hidden_chunk_meta[2] if dspark_hidden_chunk_meta else False + pd_hidden_is_last_chunk=( + pd_hidden_chunk_meta[2] if pd_hidden_chunk_meta else False ), - dspark_hidden_release_indices=( - dspark_hidden_chunk_meta[3] if dspark_hidden_chunk_meta else None + pd_hidden_release_indices=( + pd_hidden_chunk_meta[3] if pd_hidden_chunk_meta else None ), ) self._record_transfer_indices(kv_indices, state_indices) @@ -3140,7 +3140,7 @@ def send_metadata( else [None] * len(self.kv_mgr.kv_args.state_types) ) for idx, state_type in enumerate(self.kv_mgr.kv_args.state_types): - if state_type == StateType.DSPARK_HIDDEN: + if state_type == StateType.PD_HIDDEN: local_state_indices[idx] = pp_slice.get("dst_indices", []) break diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index dd19135f6898..ce651a876f11 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -91,17 +91,17 @@ def should_force_retry(req: Req) -> bool: return int.from_bytes(digest[:8], "big") < retry_prob * 2**64 -def clear_dspark_hidden_request_state(req: Req) -> None: - req.dspark_hidden_meta = None - req.dspark_hidden_src_indices = None - req.dspark_hidden_dst_indices = None - req.dspark_hidden_written = None - req.dspark_hidden_capture_layer_ids = None - req.dspark_hidden_current_src_indices = None - req.dspark_hidden_current_start = None - req.dspark_hidden_current_row_len = 0 - req.dspark_hidden_current_is_last = False - req.dspark_hidden_owner_direct_sent = False +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 def maybe_release_metadata_buffer( @@ -121,56 +121,50 @@ def maybe_release_metadata_buffer( if req.metadata_buffer_index >= 0: allocator.free(req.metadata_buffer_index) req.metadata_buffer_index = -1 - indices = getattr(req, "dspark_hidden_src_indices", None) + indices = getattr(req, "pd_hidden_src_indices", None) if indices and pd_hidden_pool is not None: sender = getattr(req, "disagg_kv_sender", None) kv_mgr = getattr(sender, "kv_mgr", None) - pop_hidden_done = getattr(kv_mgr, "pop_dspark_hidden_request_done", None) + pop_hidden_done = getattr(kv_mgr, "pop_pd_hidden_request_done", None) worker_released = pop_hidden_done is not None and pop_hidden_done( getattr(sender, "bootstrap_room", req.bootstrap_room) ) if not worker_released: pd_hidden_pool.free(indices) - clear_dspark_hidden_request_state(req) + clear_pd_hidden_request_state(req) elif not indices: - clear_dspark_hidden_request_state(req) + clear_pd_hidden_request_state(req) 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 = getattr(req, "dspark_hidden_src_indices", None) + indices = getattr(req, "pd_hidden_src_indices", None) if indices: pd_hidden_pool.free(indices) - clear_dspark_hidden_request_state(req) + clear_pd_hidden_request_state(req) def maybe_release_pd_hidden_rows_on_hidden_done( req: Req, pd_hidden_pool ) -> bool: - """Release source hidden rows after DSPARK_HIDDEN finishes, before KV success.""" - indices = getattr(req, "dspark_hidden_src_indices", None) + """Release source hidden rows after PD_HIDDEN finishes, before KV success.""" + indices = getattr(req, "pd_hidden_src_indices", None) if not indices or pd_hidden_pool is None: return False sender = getattr(req, "disagg_kv_sender", None) kv_mgr = getattr(sender, "kv_mgr", None) - pop_hidden_done = getattr(kv_mgr, "pop_dspark_hidden_request_done", None) + pop_hidden_done = getattr(kv_mgr, "pop_pd_hidden_request_done", None) if pop_hidden_done is None or not pop_hidden_done( getattr(sender, "bootstrap_room", req.bootstrap_room) ): return False - clear_dspark_hidden_request_state(req) + clear_pd_hidden_request_state(req) return True -maybe_release_dspark_hidden_rows = maybe_release_pd_hidden_rows -maybe_release_dspark_hidden_rows_on_hidden_done = ( - maybe_release_pd_hidden_rows_on_hidden_done -) - - class PrefillBootstrapQueue: """ Store the requests in bootstrapping @@ -210,7 +204,7 @@ def __init__( self.max_total_num_tokens = ( self.scheduler.tp_worker.model_runner.effective_max_total_num_tokens ) - self._last_dspark_hidden_credit_warning_time = 0.0 + self._last_pd_hidden_credit_warning_time = 0.0 self.transfer_backend = transfer_backend if envs.SGLANG_DISAGG_STAGING_BUFFER.get() and self.is_mla_backend: raise RuntimeError( @@ -357,10 +351,10 @@ def ensure_metadata_buffer(self, req: Req) -> bool: assert req.metadata_buffer_index is not None return True - def _requires_dspark_hidden_transfer(self, req: Req) -> bool: - if self.kv_manager.req_to_dspark_hidden_meta.get(req.bootstrap_room): + def _requires_pd_hidden_transfer(self, req: Req) -> bool: + if self.kv_manager.req_to_pd_hidden_meta.get(req.bootstrap_room): return True - return StateType.DSPARK_HIDDEN in self.kv_manager.kv_args.state_types + return StateType.PD_HIDDEN in self.kv_manager.kv_args.state_types def finalize_bootstrap(self, req: Req) -> bool: """Initialize the sender after bootstrap completes. @@ -374,8 +368,8 @@ def finalize_bootstrap(self, req: Req) -> bool: if decode_prefix_len is None: decode_prefix_len = req.disagg_kv_sender.pop_decode_prefix_len() req.disagg_decode_prefix_len = decode_prefix_len - dspark_meta = self.kv_manager.req_to_dspark_hidden_meta.get(req.bootstrap_room) - if dspark_meta and not self._finalize_dspark_hidden_bootstrap( + dspark_meta = self.kv_manager.req_to_pd_hidden_meta.get(req.bootstrap_room) + if dspark_meta and not self._finalize_pd_hidden_bootstrap( req, dspark_meta, decode_prefix_len ): if metadata_buffer_was_unallocated and req.metadata_buffer_index >= 0: @@ -407,7 +401,7 @@ def _probe_bootstrap_ready( if metadata_cost > metadata_credits: return None, None - dspark_meta = self.kv_manager.req_to_dspark_hidden_meta.get(req.bootstrap_room) + dspark_meta = self.kv_manager.req_to_pd_hidden_meta.get(req.bootstrap_room) if not dspark_meta: return (metadata_cost, 0), None @@ -438,11 +432,11 @@ def _probe_bootstrap_ready( return (metadata_cost, 0), None hidden_cost = 0 if plan.streaming_hidden else plan.source_window_rows - if getattr(req, "dspark_hidden_src_indices", None) is not None: + if getattr(req, "pd_hidden_src_indices", None) is not None: hidden_cost = 0 if hidden_cost > hidden_row_credits: now = time.monotonic() - if now - self._last_dspark_hidden_credit_warning_time > 30: + if now - self._last_pd_hidden_credit_warning_time > 30: logger.warning( "PD hidden pool blocked prefill bootstrap: " "rid=%s hidden_len=%d required_rows=%d free_rows=%d " @@ -454,18 +448,18 @@ def _probe_bootstrap_ready( plan.pool.size, len(self.queue), ) - self._last_dspark_hidden_credit_warning_time = now + self._last_pd_hidden_credit_warning_time = now return None, None return (metadata_cost, hidden_cost), None - def _is_dspark_hidden_credit_blocked( + def _is_pd_hidden_credit_blocked( self, req: Req, metadata_credits: int, hidden_row_credits: int ) -> bool: metadata_cost = 1 if req.metadata_buffer_index < 0 else 0 if metadata_cost > metadata_credits: return False - dspark_meta = self.kv_manager.req_to_dspark_hidden_meta.get(req.bootstrap_room) + dspark_meta = self.kv_manager.req_to_pd_hidden_meta.get(req.bootstrap_room) if not dspark_meta: return False @@ -480,7 +474,7 @@ def _is_dspark_hidden_credit_blocked( else [int(x) for x in dspark_meta.get("target_layer_ids", [])] ) ) - if not local_layer_ids or getattr(req, "dspark_hidden_src_indices", None): + if not local_layer_ids or getattr(req, "pd_hidden_src_indices", None): return False pool = getattr(self.metadata_buffers, "pd_hidden_pool", None) @@ -505,7 +499,7 @@ def stage_pp_bootstrap_consensus(self, rids: List[str]) -> List[str]: committed.append(req.rid) return committed - def _abort_dspark_hidden_bootstrap(self, req: Req, message: str) -> None: + def _abort_pd_hidden_bootstrap(self, req: Req, message: str) -> None: logger.error(message) prepare_abort(req, message, status_code=HTTPStatus.INTERNAL_SERVER_ERROR) sender = getattr(req, "disagg_kv_sender", None) @@ -517,7 +511,7 @@ def _abort_dspark_hidden_bootstrap(self, req: Req, message: str) -> None: ) sender.conclude_state = KVPoll.Failed - def _finalize_dspark_hidden_bootstrap( + def _finalize_pd_hidden_bootstrap( self, req: Req, dspark_meta: dict, decode_prefix_len: int ) -> bool: plan, error = resolve_hidden_bootstrap_plan( @@ -533,16 +527,16 @@ def _finalize_dspark_hidden_bootstrap( ), ) if error is not None: - self._abort_dspark_hidden_bootstrap(req, error) + self._abort_pd_hidden_bootstrap(req, error) return False assert plan is not None if not plan.local_layer_ids: - req.dspark_hidden_meta = dict(dspark_meta) - req.dspark_hidden_src_indices = [] - req.dspark_hidden_dst_indices = [] - req.dspark_hidden_written = [] - req.dspark_hidden_owner_direct_sent = False + 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 return True src_indices = ( @@ -556,20 +550,20 @@ def _finalize_dspark_hidden_bootstrap( f"rid={req.rid}, hidden_len={plan.hidden_len}, " f"required_rows={plan.source_window_rows}, " f"pool_size={plan.pool.size}. " - "Increase SGLANG_DSPARK_PD_HIDDEN_POOL_TOKENS or reduce the " + "Increase SGLANG_PD_HIDDEN_POOL_TOKENS or reduce the " "maximum prompt/hidden transfer length." ) - self._abort_dspark_hidden_bootstrap(req, message) + self._abort_pd_hidden_bootstrap(req, message) return False - req.dspark_hidden_capture_layer_ids = [int(x) for x in plan.local_layer_ids] - req.dspark_hidden_meta = dict(dspark_meta) - req.dspark_hidden_src_indices = src_indices - req.dspark_hidden_dst_indices = plan.dst_indices - req.dspark_hidden_written = ( + 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 = ( None if plan.streaming_hidden else [False] * plan.hidden_len ) - req.dspark_hidden_owner_direct_sent = False + req.pd_hidden_owner_direct_sent = False return True def add(self, req: Req, num_kv_heads: int) -> None: @@ -642,7 +636,7 @@ def pop_bootstrapped( indices_to_remove.add(i) failed_reqs.append(req) elif poll == KVPoll.Bootstrapping: - if self._requires_dspark_hidden_transfer(req): + if self._requires_pd_hidden_transfer(req): # PD hidden must be captured for every prefill chunk. # Do not run optimistic forward before hidden rows and # capture metadata are materialized. @@ -715,11 +709,11 @@ def get_ready_bootstrapped_rids_for_pp(self) -> Tuple[List[str], List[str]]: req, metadata_credits, hidden_row_credits ) if error is not None: - self._abort_dspark_hidden_bootstrap(req, error) + self._abort_pd_hidden_bootstrap(req, error) failed_rids.append(req.rid) continue if costs is None: - if self._is_dspark_hidden_credit_blocked( + if self._is_pd_hidden_credit_blocked( req, metadata_credits, hidden_row_credits ): break @@ -824,26 +818,26 @@ 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_dspark_hidden_capture_for_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_dspark_hidden_capture_for_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, "dspark_hidden_capture_layer_ids", None + req, "pd_hidden_capture_layer_ids", None ) if dspark_capture_layers: break if dspark_capture_layers: - batch.dspark_hidden_capture_layer_ids = [ + batch.pd_hidden_capture_layer_ids = [ int(x) for x in dspark_capture_layers ] batch.capture_hidden_mode = CaptureHiddenMode.FULL @@ -939,7 +933,7 @@ def event_loop_overlap_disagg_prefill(self: Scheduler) -> None: # Update last_batch self.last_batch = batch - def _extract_dspark_hidden_states_from_result( + def _extract_pd_hidden_states_from_result( self: Scheduler, result: GenerationBatchResult, ) -> Optional[torch.Tensor]: @@ -950,7 +944,7 @@ def _extract_dspark_hidden_states_from_result( aux_keys = sorted( key for key in proxy_tensors - if key.startswith("dspark_aux_hidden_states_") + if key.startswith("pd_aux_hidden_states_") ) if aux_keys: hidden_states = ( @@ -960,10 +954,10 @@ def _extract_dspark_hidden_states_from_result( ) return hidden_states - def _build_dspark_hidden_only_state_indices( + def _build_pd_hidden_only_state_indices( self: Scheduler, req: Req ) -> Optional[List]: - current_indices = getattr(req, "dspark_hidden_current_src_indices", None) + current_indices = getattr(req, "pd_hidden_current_src_indices", None) if current_indices is None: return None @@ -972,25 +966,25 @@ def _build_dspark_hidden_only_state_indices( ) state_indices = [] for st in state_types: - if st == StateType.DSPARK_HIDDEN: + if st == StateType.PD_HIDDEN: state_indices.append(np.asarray(current_indices, dtype=np.int32)) else: state_indices.append(None) return state_indices - def _send_dspark_hidden_only_chunk(self: Scheduler, req: Req) -> bool: - current_indices = getattr(req, "dspark_hidden_current_src_indices", None) - current_start = getattr(req, "dspark_hidden_current_start", None) - current_rows = int(getattr(req, "dspark_hidden_current_row_len", 0) or 0) + def _send_pd_hidden_only_chunk(self: Scheduler, req: Req) -> bool: + current_indices = getattr(req, "pd_hidden_current_src_indices", None) + current_start = getattr(req, "pd_hidden_current_start", None) + current_rows = int(getattr(req, "pd_hidden_current_row_len", 0) or 0) if current_indices is None or current_start is None or current_rows <= 0: return False - state_indices = self._build_dspark_hidden_only_state_indices(req) + state_indices = self._build_pd_hidden_only_state_indices(req) if state_indices is None: return False streaming_hidden = bool( - (getattr(req, "dspark_hidden_meta", None) or {}).get( + (getattr(req, "pd_hidden_meta", None) or {}).get( "streaming_hidden", False ) ) @@ -999,29 +993,29 @@ def _send_dspark_hidden_only_chunk(self: Scheduler, req: Req) -> bool: source_event.record() req.disagg_kv_sender.set_source_event(source_event) set_chunk_meta = getattr( - req.disagg_kv_sender, "set_dspark_hidden_chunk_meta", None + req.disagg_kv_sender, "set_pd_hidden_chunk_meta", None ) if set_chunk_meta is not None: set_chunk_meta( int(current_start), int(current_rows), - bool(getattr(req, "dspark_hidden_current_is_last", False)), + bool(getattr(req, "pd_hidden_current_is_last", False)), current_indices if streaming_hidden - else getattr(req, "dspark_hidden_src_indices", None), + else getattr(req, "pd_hidden_src_indices", None), ) req.disagg_kv_sender.send(np.asarray([], dtype=np.int32), state_indices) if streaming_hidden: - req.dspark_hidden_src_indices = None - req.dspark_hidden_current_src_indices = None - req.dspark_hidden_current_start = None - req.dspark_hidden_current_row_len = 0 - req.dspark_hidden_current_is_last = False - req.dspark_hidden_owner_direct_sent = True + 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 return True - def _write_dspark_hidden_rows_for_batch( + def _write_pd_hidden_rows_for_batch( self: Scheduler, batch: ScheduleBatch, result: GenerationBatchResult, @@ -1029,30 +1023,30 @@ def _write_dspark_hidden_rows_for_batch( send_owner_direct: bool = False, ) -> None: pool = getattr(self.disagg_metadata_buffers, "pd_hidden_pool", None) - hidden_states = self._extract_dspark_hidden_states_from_result(result) - needs_dspark_hidden = any( + hidden_states = self._extract_pd_hidden_states_from_result(result) + needs_pd_hidden = any( ( - getattr(req, "dspark_hidden_src_indices", None) - or getattr(req, "dspark_hidden_capture_layer_ids", None) + getattr(req, "pd_hidden_src_indices", None) + or getattr(req, "pd_hidden_capture_layer_ids", None) ) and ( send_owner_direct - or not getattr(req, "dspark_hidden_owner_direct_sent", False) + or not getattr(req, "pd_hidden_owner_direct_sent", False) ) for req in batch.reqs ) - if pool is not None and needs_dspark_hidden and hidden_states is None: + if pool is not None and needs_pd_hidden and hidden_states is None: reqs = [ ( req.rid, - getattr(req, "dspark_hidden_capture_layer_ids", None), - bool(getattr(req, "dspark_hidden_src_indices", None)), + getattr(req, "pd_hidden_capture_layer_ids", None), + bool(getattr(req, "pd_hidden_src_indices", None)), ) for req in batch.reqs ] raise RuntimeError( "PD hidden capture was required but forward output has no " - f"hidden states: batch_capture_layers={batch.dspark_hidden_capture_layer_ids}, " + f"hidden states: batch_capture_layers={batch.pd_hidden_capture_layer_ids}, " f"reqs={reqs}" ) if pool is None or hidden_states is None or batch.extend_lens is None: @@ -1077,14 +1071,14 @@ def _write_dspark_hidden_rows_for_batch( req_hidden = hidden_states[hidden_offset : hidden_offset + extend_len] hidden_offset += extend_len - meta = getattr(req, "dspark_hidden_meta", None) or {} + meta = getattr(req, "pd_hidden_meta", None) or {} streaming_hidden = bool(meta.get("streaming_hidden", False)) if ( not send_owner_direct - and getattr(req, "dspark_hidden_owner_direct_sent", False) + and getattr(req, "pd_hidden_owner_direct_sent", False) ): continue - src_indices = getattr(req, "dspark_hidden_src_indices", None) + src_indices = getattr(req, "pd_hidden_src_indices", None) if not src_indices and not streaming_hidden: continue @@ -1132,9 +1126,9 @@ def _write_dspark_hidden_rows_for_batch( else: rows = local_end - local_start write_indices = src_indices[local_start:local_end] - prev_current_start = getattr(req, "dspark_hidden_current_start", None) + prev_current_start = getattr(req, "pd_hidden_current_start", None) prev_current_row_len = int( - getattr(req, "dspark_hidden_current_row_len", 0) or 0 + getattr(req, "pd_hidden_current_row_len", 0) or 0 ) if ( prev_current_start is not None @@ -1163,20 +1157,20 @@ def _write_dspark_hidden_rows_for_batch( f"pool_rows={pool.size}. Streaming source rows are released " "only after the matching hidden chunk ACK." ) - req.dspark_hidden_src_indices = write_indices + req.pd_hidden_src_indices = write_indices pool.write( write_indices, req_hidden_to_write[chunk_local_start:chunk_local_end], ) - req.dspark_hidden_current_start = write_start - req.dspark_hidden_current_row_len = rows - req.dspark_hidden_current_src_indices = write_indices - req.dspark_hidden_current_is_last = write_end >= hidden_start + hidden_len - written = getattr(req, "dspark_hidden_written", None) + 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 = getattr(req, "pd_hidden_written", None) if written is not None: written[local_start:local_end] = [True] * rows if send_owner_direct: - self._send_dspark_hidden_only_chunk(req) + self._send_pd_hidden_only_chunk(req) def send_dspark_owner_direct_hidden_for_batch( self: Scheduler, @@ -1186,7 +1180,7 @@ def send_dspark_owner_direct_hidden_for_batch( capture_reqs = [ req for req in batch.reqs - if getattr(req, "dspark_hidden_capture_layer_ids", None) + if getattr(req, "pd_hidden_capture_layer_ids", None) ] if not capture_reqs: return False @@ -1194,14 +1188,14 @@ def send_dspark_owner_direct_hidden_for_batch( return False if not all( bool( - (getattr(req, "dspark_hidden_meta", None) or {}).get( + (getattr(req, "pd_hidden_meta", None) or {}).get( "streaming_hidden", False ) ) for req in capture_reqs ): return False - self._write_dspark_hidden_rows_for_batch( + self._write_pd_hidden_rows_for_batch( batch, result, send_owner_direct=True ) return True @@ -1247,7 +1241,7 @@ def process_batch_result_disagg_prefill( batch=batch, logits_output=logits_output, ) - self._write_dspark_hidden_rows_for_batch(batch, result) + self._write_pd_hidden_rows_for_batch(batch, result) def advance_logprob_pt(i: int, req: Req) -> None: nonlocal logprob_pt @@ -1626,7 +1620,7 @@ def process_prefill_chunk( if is_aborted(req): # bootstrap failed self.chunked_req = None - elif self.disagg_prefill_bootstrap_queue._requires_dspark_hidden_transfer( + elif self.disagg_prefill_bootstrap_queue._requires_pd_hidden_transfer( req ): self.chunked_req = None @@ -1708,25 +1702,25 @@ def send_kv_chunk( ) return True - current_dspark_hidden_src_indices = getattr( - req, "dspark_hidden_current_src_indices", None + current_pd_hidden_src_indices = getattr( + req, "pd_hidden_current_src_indices", None ) - current_dspark_hidden_start = getattr(req, "dspark_hidden_current_start", None) - current_dspark_hidden_row_len = int( - getattr(req, "dspark_hidden_current_row_len", 0) or 0 + current_pd_hidden_start = getattr(req, "pd_hidden_current_start", None) + current_pd_hidden_row_len = int( + getattr(req, "pd_hidden_current_row_len", 0) or 0 ) - has_current_dspark_hidden = ( - current_dspark_hidden_src_indices is not None - and current_dspark_hidden_row_len > 0 + has_current_pd_hidden = ( + current_pd_hidden_src_indices is not None + and current_pd_hidden_row_len > 0 ) - streaming_dspark_hidden = bool( - (getattr(req, "dspark_hidden_meta", None) or {}).get( + streaming_pd_hidden = bool( + (getattr(req, "pd_hidden_meta", None) or {}).get( "streaming_hidden", False ) ) state_indices: Optional[List] = None - if last_chunk or has_current_dspark_hidden: + if last_chunk or has_current_pd_hidden: if last_chunk: self.disagg_metadata_buffers.set_buf(req) @@ -1795,23 +1789,23 @@ def _c128_state_payload(): ring_size=ring_size, ) - def _dspark_hidden_payload(): - if getattr(req, "dspark_hidden_owner_direct_sent", False): + def _pd_hidden_payload(): + if getattr(req, "pd_hidden_owner_direct_sent", False): return [] - src_indices = getattr(req, "dspark_hidden_src_indices", None) + src_indices = getattr(req, "pd_hidden_src_indices", None) if ( src_indices is None - and getattr(req, "dspark_hidden_capture_layer_ids", None) + and getattr(req, "pd_hidden_capture_layer_ids", None) ): raise RuntimeError( "PD hidden row pool was not materialized before transfer: " f"rid={req.rid}" ) - if has_current_dspark_hidden: - return np.asarray(current_dspark_hidden_src_indices, dtype=np.int32) + if has_current_pd_hidden: + return np.asarray(current_pd_hidden_src_indices, dtype=np.int32) if not src_indices: return [] - written = getattr(req, "dspark_hidden_written", None) + written = getattr(req, "pd_hidden_written", None) if written is not None and not all(written): missing = [i for i, ok in enumerate(written) if not ok][:8] raise RuntimeError( @@ -1839,8 +1833,8 @@ def _dspark_hidden_payload(): state_indices.append(_swa_ring_payload()) elif st == StateType.C128_STATE: state_indices.append(_c128_state_payload()) - elif st == StateType.DSPARK_HIDDEN: - state_indices.append(_dspark_hidden_payload()) + elif st == StateType.PD_HIDDEN: + state_indices.append(_pd_hidden_payload()) else: state_indices.append(None) @@ -1851,34 +1845,34 @@ def _dspark_hidden_payload(): should_send_kv_chunk = req.disagg_kv_sender.should_send_kv_chunk( len(page_indices), last_chunk ) - if not should_send_kv_chunk and not has_current_dspark_hidden: + if not should_send_kv_chunk and not has_current_pd_hidden: return True if ( - has_current_dspark_hidden + has_current_pd_hidden and hasattr(req.disagg_kv_sender, "set_source_event") ): source_event = self.device_module.Event() source_event.record() req.disagg_kv_sender.set_source_event(source_event) set_chunk_meta = getattr( - req.disagg_kv_sender, "set_dspark_hidden_chunk_meta", None + req.disagg_kv_sender, "set_pd_hidden_chunk_meta", None ) if set_chunk_meta is not None: set_chunk_meta( - int(current_dspark_hidden_start), - int(current_dspark_hidden_row_len), - bool(getattr(req, "dspark_hidden_current_is_last", False)), - current_dspark_hidden_src_indices - if streaming_dspark_hidden - else getattr(req, "dspark_hidden_src_indices", None), + int(current_pd_hidden_start), + int(current_pd_hidden_row_len), + bool(getattr(req, "pd_hidden_current_is_last", False)), + current_pd_hidden_src_indices + if streaming_pd_hidden + else getattr(req, "pd_hidden_src_indices", None), ) req.disagg_kv_sender.send(page_indices, state_indices) - if has_current_dspark_hidden and streaming_dspark_hidden: - req.dspark_hidden_src_indices = None - req.dspark_hidden_current_src_indices = None - req.dspark_hidden_current_start = None - req.dspark_hidden_current_row_len = 0 - req.dspark_hidden_current_is_last = False + 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 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 cff9268e798f..9624f2574be0 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -242,7 +242,7 @@ def free(self, free_index: int): self.free_slots.append(free_index) -class DSparkHiddenRowPool: +class PDHiddenRowPool: """Compact row pool for PD hidden-state transfer.""" def __init__( @@ -500,9 +500,6 @@ def trim_dynamic_dst( return new_dynamic_dst -DSparkHiddenTransferPlan = PDHiddenTransferPlan - - class MetadataBuffers: def __init__( self, @@ -519,9 +516,9 @@ def __init__( ): self.custom_mem_pool = custom_mem_pool self.output_dsa_topk_indices_dim = output_dsa_topk_indices_dim - self.pd_hidden_pool: Optional[DSparkHiddenRowPool] = None + self.pd_hidden_pool: Optional[PDHiddenRowPool] = None if pd_hidden_pool_size > 0 and pd_hidden_size > 0: - self.pd_hidden_pool = DSparkHiddenRowPool( + self.pd_hidden_pool = PDHiddenRowPool( pd_hidden_pool_size, pd_hidden_size, hidden_states_dtype, @@ -795,9 +792,9 @@ def ensure_pd_hidden_pool( hidden_size: int, dtype: torch.dtype, device: str = "cpu", - ) -> DSparkHiddenRowPool: + ) -> PDHiddenRowPool: if self.pd_hidden_pool is None: - self.pd_hidden_pool = DSparkHiddenRowPool( + self.pd_hidden_pool = PDHiddenRowPool( size=size, hidden_size=hidden_size, dtype=dtype, @@ -1199,7 +1196,7 @@ def setup_state_kv_args( draft_token_to_kv_pool=None, total_kv_layers: int = None, req_to_token_pool=None, - pd_hidden_pool: Optional[DSparkHiddenRowPool] = None, + pd_hidden_pool: Optional[PDHiddenRowPool] = None, ) -> None: """Populate ``kv_args`` state-buffer fields from the given pool. Shared by prefill and decode bootstrap paths so the state_type dispatch @@ -1319,7 +1316,7 @@ def setup_state_kv_args( if data_ptrs: append_state_component( kv_args, - StateType.DSPARK_HIDDEN, + StateType.PD_HIDDEN, data_ptrs, data_lens, item_lens, diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 5916a9720a3f..8b6a0cffcb53 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -291,7 +291,7 @@ class Envs: SGLANG_DSPARK_OPT_MARKOV_W2_BF16 = EnvBool(True) SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD = EnvBool(True) SGLANG_DSPARK_ENABLE_MULTI_STREAM = EnvBool(True) - SGLANG_DSPARK_PD_HIDDEN_RECV_POOL_TOKENS = EnvInt(-1) + SGLANG_PD_HIDDEN_RECV_POOL_TOKENS = EnvInt(-1) SGLANG_DEBUG_REVERT_PR = EnvInt(0) SGLANG_PHASE_CHECKER_DEBUG = EnvBool(False) SGLANG_TEST_REQUEST_TIME_STATS = EnvBool(False) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index f1b921797eea..882ba48413d1 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1977,7 +1977,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): seq_lens_cpu_cache: torch.Tensor = None capture_hidden_mode: Optional[CaptureHiddenMode] = None return_hidden_states_before_norm: bool = False - dspark_hidden_capture_layer_ids: Optional[List[int]] = None + pd_hidden_capture_layer_ids: Optional[List[int]] = None @classmethod def init_new( @@ -3033,9 +3033,9 @@ def copy(self): prefill_stats=self.prefill_stats, fpm_start_time=self.fpm_start_time, forward_iter=self.forward_iter, - dspark_hidden_capture_layer_ids=( - self.dspark_hidden_capture_layer_ids[:] - if self.dspark_hidden_capture_layer_ids is not None + pd_hidden_capture_layer_ids=( + self.pd_hidden_capture_layer_ids[:] + if self.pd_hidden_capture_layer_ids is not None else None ), ) diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 0e9ba6c63181..4fd1d6c29ebf 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -259,7 +259,7 @@ 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_dspark_hidden_capture_for_batch(batch) + self._prepare_pd_hidden_capture_for_batch(batch) self.mbs[mb_id] = batch self.running_mbs[mb_id] = self.running_batch @@ -1016,25 +1016,25 @@ def _pp_prepare_tensor_dict( **logprob_dict, } if ( - batch.dspark_hidden_capture_layer_ids - and not self._pp_should_owner_direct_dspark_hidden(batch) + batch.pd_hidden_capture_layer_ids + 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 ): - tensor_dict["dspark_aux_hidden_states_0"] = result.logits_output.hidden_states + tensor_dict["pd_aux_hidden_states_0"] = result.logits_output.hidden_states return tensor_dict - def _pp_should_owner_direct_dspark_hidden( + def _pp_should_owner_direct_pd_hidden( self: Scheduler, batch: ScheduleBatch ) -> bool: if not hasattr(self, "disagg_prefill_bootstrap_queue"): return False - if not batch or not batch.dspark_hidden_capture_layer_ids: + if not batch or not batch.pd_hidden_capture_layer_ids: return False capture_reqs = [ req for req in batch.reqs - if getattr(req, "dspark_hidden_capture_layer_ids", None) + if getattr(req, "pd_hidden_capture_layer_ids", None) ] if not capture_reqs: return False @@ -1042,14 +1042,14 @@ def _pp_should_owner_direct_dspark_hidden( return False return all( bool( - (getattr(req, "dspark_hidden_meta", None) or {}).get( + (getattr(req, "pd_hidden_meta", None) or {}).get( "streaming_hidden", False ) ) for req in capture_reqs ) - def _pp_strip_dspark_aux_hidden_from_proxy( + def _pp_strip_pd_aux_hidden_from_proxy( self: Scheduler, result: GenerationBatchResult ) -> None: proxy = result.pp_hidden_states_proxy_tensors @@ -1057,7 +1057,7 @@ def _pp_strip_dspark_aux_hidden_from_proxy( return tensors = proxy.tensors aux_keys = [ - key for key in tensors if key.startswith("dspark_aux_hidden_states_") + key for key in tensors if key.startswith("pd_aux_hidden_states_") ] for key in aux_keys: tensors.pop(key, None) @@ -1067,7 +1067,7 @@ def _pp_maybe_send_dspark_owner_direct_hidden( batch: ScheduleBatch, result: GenerationBatchResult, ) -> None: - if not self._pp_should_owner_direct_dspark_hidden(batch): + if not self._pp_should_owner_direct_pd_hidden(batch): return send_owner_direct = getattr( self, "send_dspark_owner_direct_hidden_for_batch", None @@ -1075,7 +1075,7 @@ def _pp_maybe_send_dspark_owner_direct_hidden( if send_owner_direct is None: return if send_owner_direct(batch, result): - self._pp_strip_dspark_aux_hidden_from_proxy(result) + self._pp_strip_pd_aux_hidden_from_proxy(result) def _pp_send_dict_to_next_stage( self: Scheduler, @@ -1206,17 +1206,17 @@ def _pp_prep_batch_result( self.future_map.stash( batch.req_pool_indices, RelayPayload(bonus_tokens=next_token_ids) ) - dspark_aux_hidden = { + pd_aux_hidden = { key: value for key, value in pp_outputs.tensors.items() - if key.startswith("dspark_aux_hidden_states_") + if key.startswith("pd_aux_hidden_states_") } batch.input_ids = None output_result = GenerationBatchResult( logits_output=logits_output, pp_hidden_states_proxy_tensors=( - PPProxyTensors(dspark_aux_hidden) if dspark_aux_hidden else None + PPProxyTensors(pd_aux_hidden) if pd_aux_hidden else None ), next_token_ids=pp_outputs["next_token_ids"], extend_input_len_per_req=extend_input_len_per_req, diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 8e0e436d5965..981b39ff55f7 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -443,7 +443,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): reuse_dsa_topk_indices: Optional[bool] = False # DeepSeek-V4 DSpark PD: per-prefill-batch target aux hidden layers to capture. - dspark_hidden_capture_layer_ids: Optional[List[int]] = None + pd_hidden_capture_layer_ids: Optional[List[int]] = None minimax_m3_precached_sparse_layers: Optional[Set[int]] = None @@ -708,7 +708,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, - dspark_hidden_capture_layer_ids=batch.dspark_hidden_capture_layer_ids, + pd_hidden_capture_layer_ids=batch.pd_hidden_capture_layer_ids, tbo_split_seq_index=batch.tbo_split_seq_index, # Host-side metadata top_logprobs_nums=batch.top_logprobs_nums, diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index cb2e5e429159..86e6179ec004 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -2252,19 +2252,19 @@ def forward( input_embeds: Optional[torch.Tensor], pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> Union[torch.Tensor, PPProxyTensors]: - incoming_dspark_aux_hidden_states: List[torch.Tensor] = [] + incoming_pd_aux_hidden_states: List[torch.Tensor] = [] if self.pp_group.is_first_rank: hidden_states = self.embed_tokens(input_ids) hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1) else: assert pp_proxy_tensors is not None hidden_states = pp_proxy_tensors["hidden_states"] - incoming_dspark_aux_hidden_states = [ + incoming_pd_aux_hidden_states = [ pp_proxy_tensors[key] for key in sorted( key for key in pp_proxy_tensors.tensors - if key.startswith("dspark_aux_hidden_states_") + if key.startswith("pd_aux_hidden_states_") ) ] if hidden_states.shape[0] != positions.shape[0]: @@ -2307,7 +2307,7 @@ def forward( delattr(forward_batch, _attr) dspark_layers_to_capture = getattr( - forward_batch, "dspark_hidden_capture_layer_ids", None + forward_batch, "pd_hidden_capture_layer_ids", None ) if dspark_layers_to_capture is None: dspark_layers_to_capture = self.dspark_layers_to_capture @@ -2318,8 +2318,8 @@ def forward( "DeepSeek-V4 prefill context parallelism (attn_cp_size > 1). Disable one " "of them: DSpark static-verify is CP-off for v1." ) - dspark_aux_hidden_states: List[torch.Tensor] = list( - incoming_dspark_aux_hidden_states + pd_aux_hidden_states: List[torch.Tensor] = list( + incoming_pd_aux_hidden_states ) # DSpark aux capture needs the per-layer eager loop (TBO's overlapped # execution cannot expose per-layer completed hidden states), so skip @@ -2362,7 +2362,7 @@ def forward( ) else: completed = hidden_states - dspark_aux_hidden_states.append(completed.mean(dim=1)) + pd_aux_hidden_states.append(completed.mean(dim=1)) if use_fused and last_layer is not None: hidden_states = last_layer.hc_post( hidden_states, prev_residual, prev_post, prev_comb @@ -2381,8 +2381,8 @@ def forward( # Flatten 3D mHC tensor for PP IPC. proxy_tensors = {"hidden_states": hidden_states.flatten(1)} if capture_dspark: - for idx, aux_hidden in enumerate(dspark_aux_hidden_states): - proxy_tensors[f"dspark_aux_hidden_states_{idx}"] = ( + for idx, aux_hidden in enumerate(pd_aux_hidden_states): + proxy_tensors[f"pd_aux_hidden_states_{idx}"] = ( aux_hidden.flatten(1) if aux_hidden.ndim == 3 else aux_hidden ) return PPProxyTensors(proxy_tensors) @@ -2395,7 +2395,7 @@ def forward( hidden_states = self.norm(hidden_states) if capture_dspark: - return (hidden_states, pre_hc_head), dspark_aux_hidden_states + return (hidden_states, pre_hc_head), pd_aux_hidden_states return hidden_states, pre_hc_head @@ -2550,14 +2550,14 @@ def forward( return hidden_states aux_hidden_states = None - dspark_aux_hidden_states = None - has_dspark_hidden_capture = ( - getattr(forward_batch, "dspark_hidden_capture_layer_ids", None) is not None + pd_aux_hidden_states = None + has_pd_hidden_capture = ( + getattr(forward_batch, "pd_hidden_capture_layer_ids", None) is not None ) - if has_dspark_hidden_capture: - hidden_states, dspark_aux_hidden_states = hidden_states + if has_pd_hidden_capture: + hidden_states, pd_aux_hidden_states = hidden_states if self.capture_aux_hidden_states: - aux_hidden_states = dspark_aux_hidden_states + aux_hidden_states = pd_aux_hidden_states elif self.capture_aux_hidden_states: hidden_states, aux_hidden_states = hidden_states hidden_states, pre_hc_head = hidden_states @@ -2573,12 +2573,12 @@ def forward( ), ) if ( - has_dspark_hidden_capture - and dspark_aux_hidden_states + has_pd_hidden_capture + and pd_aux_hidden_states and logits_output.hidden_states is None ): flattened_aux_hidden_states = [ - x.flatten(1) if x.ndim == 3 else x for x in dspark_aux_hidden_states + x.flatten(1) if x.ndim == 3 else x for x in pd_aux_hidden_states ] logits_output.hidden_states = ( flattened_aux_hidden_states[0] diff --git a/python/sglang/srt/speculative/dspark_components/dspark_disaggregation.py b/python/sglang/srt/speculative/dspark_components/dspark_disaggregation.py index 01696742cdb8..d0b23b321702 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_disaggregation.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_disaggregation.py @@ -33,7 +33,7 @@ def metadata_buffer_kwargs(self) -> dict: @dataclass -class DSparkHiddenBootstrapPlan: +class PDHiddenBootstrapPlan: hidden_start: int hidden_len: int streaming_hidden: bool @@ -109,11 +109,11 @@ def resolve_disagg_metadata_config( int(server_args.max_prefill_buffer_tokens() or 0), int(max_prefill_tokens or 0), ) - pool_env_value = os.getenv("SGLANG_DSPARK_PD_HIDDEN_POOL_TOKENS") + pool_env_value = os.getenv("SGLANG_PD_HIDDEN_POOL_TOKENS") if mode_value == "decode": - pool_env_value = os.getenv("SGLANG_DSPARK_PD_HIDDEN_RECV_POOL_TOKENS") + pool_env_value = os.getenv("SGLANG_PD_HIDDEN_RECV_POOL_TOKENS") if pool_env_value is None: - pool_env_value = os.getenv("SGLANG_DSPARK_PD_HIDDEN_POOL_TOKENS") + pool_env_value = os.getenv("SGLANG_PD_HIDDEN_POOL_TOKENS") hidden_pool_size = max( 0, int(pool_env_value if pool_env_value is not None else str(default_pool_rows)), @@ -148,7 +148,7 @@ def resolve_hidden_bootstrap_plan( model_runner: Any, metadata_buffers: Any, prefill_radix_enabled: bool, -) -> Tuple[Optional[DSparkHiddenBootstrapPlan], Optional[str]]: +) -> Tuple[Optional[PDHiddenBootstrapPlan], Optional[str]]: hidden_start = int(metadata.get("hidden_start", 0)) hidden_len = int(metadata.get("hidden_len", len(req.origin_input_ids))) if hidden_start != int(decode_prefix_len): @@ -190,7 +190,7 @@ def resolve_hidden_bootstrap_plan( ) if not local_layer_ids: return ( - DSparkHiddenBootstrapPlan( + PDHiddenBootstrapPlan( hidden_start=hidden_start, hidden_len=hidden_len, streaming_hidden=bool(metadata.get("streaming_hidden", False)), @@ -270,7 +270,7 @@ def resolve_hidden_bootstrap_plan( ) return ( - DSparkHiddenBootstrapPlan( + PDHiddenBootstrapPlan( hidden_start=hidden_start, hidden_len=hidden_len, streaming_hidden=streaming_hidden, diff --git a/test/registered/unit/disaggregation/test_dspark_hidden_state.py b/test/registered/unit/disaggregation/test_pd_hidden_state.py similarity index 83% rename from test/registered/unit/disaggregation/test_dspark_hidden_state.py rename to test/registered/unit/disaggregation/test_pd_hidden_state.py index 4349a6c7dbc4..06e05c11cfab 100644 --- a/test/registered/unit/disaggregation/test_dspark_hidden_state.py +++ b/test/registered/unit/disaggregation/test_pd_hidden_state.py @@ -3,14 +3,14 @@ import torch from sglang.srt.disaggregation.common.utils import ( - DSparkHiddenChunk, - DSparkHiddenRequestState, + PDHiddenChunk, + PDHiddenRequestState, ) -from sglang.srt.disaggregation.utils import DSparkHiddenRowPool +from sglang.srt.disaggregation.utils import PDHiddenRowPool -def _chunk(start: int, rows: int, is_last: bool = False) -> DSparkHiddenChunk: - return DSparkHiddenChunk( +def _chunk(start: int, rows: int, is_last: bool = False) -> PDHiddenChunk: + return PDHiddenChunk( room=1, prefill_rank=0, hidden_start=start, @@ -20,9 +20,9 @@ def _chunk(start: int, rows: int, is_last: bool = False) -> DSparkHiddenChunk: ) -class TestDSparkHiddenRequestState(unittest.TestCase): +class TestPDHiddenRequestState(unittest.TestCase): def test_disabled_state_is_done_for_hidden_but_not_kv(self): - state = DSparkHiddenRequestState.disabled() + state = PDHiddenRequestState.disabled() self.assertFalse(state.enabled) self.assertFalse(state.streaming) @@ -34,7 +34,7 @@ def test_disabled_state_is_done_for_hidden_but_not_kv(self): self.assertTrue(state.request_done()) def test_full_state_waits_only_for_kv_done(self): - state = DSparkHiddenRequestState.full(2, 6) + state = PDHiddenRequestState.full(2, 6) self.assertTrue(state.enabled) self.assertFalse(state.streaming) @@ -48,7 +48,7 @@ def test_full_state_waits_only_for_kv_done(self): self.assertTrue(state.request_done()) def test_streaming_hidden_done_is_separate_from_request_done(self): - state = DSparkHiddenRequestState.streaming_state(0, 8) + state = PDHiddenRequestState.streaming_state(0, 8) self.assertEqual(state.accept_chunk(_chunk(0, 4)), "accepted") self.assertFalse(state.hidden_request_done()) @@ -63,7 +63,7 @@ def test_streaming_hidden_done_is_separate_from_request_done(self): self.assertTrue(state.request_done()) def test_streaming_hidden_completion_can_wait_for_ack(self): - state = DSparkHiddenRequestState.streaming_state(0, 8) + state = PDHiddenRequestState.streaming_state(0, 8) self.assertEqual(state.accept_chunk(_chunk(0, 4)), "accepted") self.assertEqual( @@ -77,26 +77,26 @@ def test_streaming_hidden_completion_can_wait_for_ack(self): self.assertTrue(state.hidden_request_done()) def test_streaming_hidden_rejects_future_and_stale_chunks(self): - state = DSparkHiddenRequestState.streaming_state(0, 8) + state = PDHiddenRequestState.streaming_state(0, 8) self.assertEqual(state.accept_chunk(_chunk(4, 4)), "future") self.assertEqual(state.accept_chunk(_chunk(0, 4)), "accepted") self.assertEqual(state.accept_chunk(_chunk(0, 4)), "stale") def test_streaming_hidden_last_chunk_must_end_at_expected_offset(self): - state = DSparkHiddenRequestState.streaming_state(0, 8) + state = PDHiddenRequestState.streaming_state(0, 8) with self.assertRaisesRegex(RuntimeError, "unexpected offset"): state.accept_chunk(_chunk(0, 4, is_last=True)) def test_streaming_hidden_chunk_cannot_exceed_expected_range(self): - state = DSparkHiddenRequestState.streaming_state(0, 8) + state = PDHiddenRequestState.streaming_state(0, 8) with self.assertRaisesRegex(RuntimeError, "exceeds request range"): state.accept_chunk(_chunk(0, 9)) def test_streaming_hidden_reset_returns_to_disabled_state(self): - state = DSparkHiddenRequestState.streaming_state(0, 8) + state = PDHiddenRequestState.streaming_state(0, 8) self.assertEqual(state.accept_chunk(_chunk(0, 8, is_last=True)), "accepted") state.mark_kv_done() self.assertTrue(state.request_done()) @@ -113,7 +113,7 @@ def test_streaming_hidden_reset_returns_to_disabled_state(self): self.assertFalse(state.request_done()) def test_hidden_chunk_descriptor_keeps_ack_endpoint_metadata(self): - chunk = DSparkHiddenChunk( + chunk = PDHiddenChunk( room=3, prefill_rank=7, hidden_start=16, @@ -134,9 +134,9 @@ def test_hidden_chunk_descriptor_keeps_ack_endpoint_metadata(self): self.assertEqual(chunk.ack_port, 12345) -class TestDSparkHiddenRowPool(unittest.TestCase): +class TestPDHiddenRowPool(unittest.TestCase): def test_alloc_prefers_contiguous_rows_and_merges_frees(self): - pool = DSparkHiddenRowPool(8, 1, torch.float32) + pool = PDHiddenRowPool(8, 1, torch.float32) self.assertEqual(pool.alloc(3), [0, 1, 2]) self.assertEqual(pool.alloc(2), [3, 4]) @@ -146,7 +146,7 @@ def test_alloc_prefers_contiguous_rows_and_merges_frees(self): self.assertEqual(pool.available_size(), 3) def test_alloc_falls_back_to_fragmented_rows_without_global_sorting(self): - pool = DSparkHiddenRowPool(8, 1, torch.float32) + pool = PDHiddenRowPool(8, 1, torch.float32) self.assertEqual(pool.alloc(8), list(range(8))) pool.free([0, 1, 4, 5, 7]) @@ -155,7 +155,7 @@ def test_alloc_falls_back_to_fragmented_rows_without_global_sorting(self): self.assertEqual(pool.available_size(), 2) def test_free_ignores_duplicate_and_already_free_rows(self): - pool = DSparkHiddenRowPool(4, 1, torch.float32) + pool = PDHiddenRowPool(4, 1, torch.float32) allocated = pool.alloc(2) self.assertEqual(allocated, [0, 1])