Skip to content
33 changes: 33 additions & 0 deletions python/sglang/srt/disaggregation/common/conn.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,9 @@ def __init__(
self.is_hybrid_mla_backend = getattr(args, "is_hybrid_mla_backend", False)
self.disaggregation_mode = disaggregation_mode
self.server_args = server_args
self.enable_deferred_decode_kv_release = (
envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get()
)
# for p/d multi node infer
self.bootstrap_host = server_args.host
self.bootstrap_port = server_args.disaggregation_bootstrap_port
Expand Down Expand Up @@ -251,6 +254,10 @@ def __init__(
self.session_pool_lock = threading.Lock()
self.addr_to_rooms_tracker: Dict[str, Set[int]] = defaultdict(set)
self.prefill_response_tracker: Dict[int, Set[int]] = defaultdict(set)
# Deferred KV release: room -> prefill ranks that acked their transfer
# drained. Entry exists only while the room is held, so a stale/late
# ack for a reused bootstrap_room is dropped.
self._deferred_abort_ack_tracker: Dict[int, Set[int]] = {}
# Heartbeat interval should be at least 2 seconds
self.heartbeat_interval = max(
envs.SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL.get(), 2.0
Expand Down Expand Up @@ -335,6 +342,28 @@ def record_failure(self, bootstrap_room: int, failure_reason: str):
with self.failure_lock:
self.failure_records[bootstrap_room] = failure_reason

def register_deferred_abort_room(self, bootstrap_room: int) -> None:
"""Arm drain-ack accounting for a held room; a fresh set wipes stale acks
from a prior request that reused this bootstrap_room."""
self._deferred_abort_ack_tracker[bootstrap_room] = set()

def note_abort_ack(self, bootstrap_room: int, prefill_rank: int) -> None:
"""Record a prefill rank's drain ack (decode receiver thread). Only counts
while the room is held; grabs the set by reference to avoid racing clear."""
acks = self._deferred_abort_ack_tracker.get(bootstrap_room)
if acks is not None:
acks.add(prefill_rank)

def is_abort_release_safe(self, bootstrap_room: int, required_acks: int) -> bool:
"""True once every prefill rank that could still write these pages has acked."""
return (
len(self._deferred_abort_ack_tracker.get(bootstrap_room, ()))
>= required_acks
)

def clear_deferred_abort_state(self, bootstrap_room: int) -> None:
self._deferred_abort_ack_tracker.pop(bootstrap_room, None)

def get_kv_replica_factor(self) -> int:
if self._kv_replica_factor is None:
logger.warning_once(
Expand Down Expand Up @@ -1237,6 +1266,10 @@ def clear(self) -> None:
self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, None)
if hasattr(self.kv_mgr, "transfer_infos"):
self.kv_mgr.transfer_infos.pop(self.bootstrap_room, None)
if hasattr(self.kv_mgr, "_deferred_ack_targets"):
# Drop a held ack target if the room concluded without draining
# (e.g. aborted before any chunk enqueued); else it leaks on prefill.
self.kv_mgr._deferred_ack_targets.pop(self.bootstrap_room, None)

def abort(self):
self.kv_mgr.record_failure(
Expand Down
99 changes: 94 additions & 5 deletions python/sglang/srt/disaggregation/decode.py
Original file line number Diff line number Diff line change
Expand Up @@ -1816,6 +1816,15 @@ def __init__(
self.spec_algorithm = scheduler.spec_algorithm
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
self.staging_handler = None
self.enable_deferred_kv_release = (
envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get()
)
self.deferred_kv_release_timeout = (
envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE_TIMEOUT.get()
)
# Aborted-mid-transfer requests whose KV pages/slot are held until drained
# or timed out. Entries: (decode_req, deadline, metadata_idx, required_acks).
self._deferred_releases: List[Tuple[DecodeRequest, float, int, int]] = []

def add(self, decode_req: DecodeRequest) -> None:
self.queue.append(decode_req)
Expand Down Expand Up @@ -2035,6 +2044,9 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req

transferred_reqs = []
indices_to_remove = set()
# Queue-removed but held for deferred release; excluded from the metadata
# teardown below.
deferred_indices = set()
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
Expand Down Expand Up @@ -2072,11 +2084,23 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req
)
if self.scheduler.enable_hisparse:
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)
decode_req.kv_receiver.clear()
decode_req.kv_receiver = None
indices_to_remove.add(i)
if (
self.enable_deferred_kv_release
and decode_req.kv_receiver.abort_notified
):
# Decode-initiated abort: a prefill write may still target
# these pages, so hold them until the drain ack or timeout.
# (A prefill-initiated failure has already stopped writing ->
# immediate release below.)
self._defer_release(decode_req)
deferred_indices.add(i)
indices_to_remove.add(i)
else:
# 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)
decode_req.kv_receiver.clear()
decode_req.kv_receiver = None
indices_to_remove.add(i)
if self.scheduler.metrics_reporter.enable_metrics:
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
continue
Expand Down Expand Up @@ -2114,6 +2138,9 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req
raise ValueError(f"Unexpected poll case: {poll}")

for i in indices_to_remove:
if i in deferred_indices:
# Held for deferred release; metadata buffer freed at resolve time.
continue
if self.enable_staging and self.staging_handler.is_staging_room(
self.queue[i].req.bootstrap_room
):
Expand All @@ -2133,9 +2160,67 @@ def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req

return transferred_reqs

def _defer_release(self, decode_req: DecodeRequest) -> None:
deadline = time.monotonic() + self.deferred_kv_release_timeout
# Require an ack from every notified prefill rank (dummy-proof). Snapshot
# now -- the receiver may be cleared by resolve time.
required_acks = len(decode_req.kv_receiver.bootstrap_infos)
self._deferred_releases.append(
(decode_req, deadline, decode_req.metadata_buffer_index, required_acks)
)

def _do_release(self, decode_req: DecodeRequest, idx: int) -> None:
room = decode_req.req.bootstrap_room
if self.enable_staging and self.staging_handler.is_staging_room(room):
self.staging_handler.unregister_decode_req(room)
# 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.metadata_buffers.bootstrap_room[idx] = 0
self.req_to_metadata_buffer_idx_allocator.free(idx)
decode_req.kv_receiver.kv_mgr.clear_deferred_abort_state(room)
decode_req.kv_receiver.clear()
decode_req.kv_receiver = None

def has_pending_deferred_releases(self) -> bool:
return bool(self._deferred_releases)

def resolve_deferred_releases(self) -> None:
"""Release held requests once every prefill rank acks the drain, or the
hold times out."""
if not self._deferred_releases:
return
now = time.monotonic()
still_held = []
to_release = []
for decode_req, deadline, idx, required_acks in self._deferred_releases:
room = decode_req.req.bootstrap_room
kv_mgr = decode_req.kv_receiver.kv_mgr
drained = kv_mgr.is_abort_release_safe(room, required_acks)
if not drained and now < deadline:
still_held.append((decode_req, deadline, idx, required_acks))
else:
to_release.append((decode_req, idx, room, drained))
# Commit the survivors before releasing so a _do_release exception can't
# leave a released entry in the list (double-free / None receiver on retry).
self._deferred_releases = still_held
for decode_req, idx, room, drained in to_release:
if not drained:
logger.warning(
f"Deferred KV release for room {room} timed out after "
f"{self.deferred_kv_release_timeout}s without a full drain "
f"ack from prefill; releasing anyway."
)
try:
self._do_release(decode_req, idx)
except Exception:
# Isolate a failed release so the rest still run; entry already dropped.
logger.exception(f"Deferred KV release failed for room {room}")

def release_memory_occupation(self):
"""Clean up in-flight transfers before releasing GPU memory."""
self.queue.clear()
# Pool is being torn down; drop held entries without per-request release.
self._deferred_releases.clear()

def resume_memory_occupation(self):
"""Queues are already cleared on release; new transfers can be accepted."""
Expand Down Expand Up @@ -2354,6 +2439,10 @@ def process_decode_queue(self: Scheduler):
if get_disagg().disaggregation_decode_enable_offload_kvcache:
self.decode_offload_manager.check_offload_progress()

# Resolve held releases every iteration (before the retraction/polling
# gates below) so their timeouts fire under memory pressure.
self.disagg_decode_transfer_queue.resolve_deferred_releases()

# try to resume retracted requests if there are enough space for another `num_reserved_decode_tokens` decode steps
resumed_reqs = self.disagg_decode_prealloc_queue.resume_retracted_reqs()
self.waiting_queue.extend(resumed_reqs)
Expand Down
Loading
Loading