[PD] Deferred decode-side KV release for aborts mid-transfer (#35049)

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Shangming Cai
2026-08-18 23:56:33 +08:00
committed by GitHub
co-authored by Claude Opus 4.8
parent c60952d933
commit 97dedd1ce9
8 changed files with 507 additions and 34 deletions
@@ -160,6 +160,9 @@ class CommonKVManager(BaseKVManager):
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
@@ -251,6 +254,10 @@ class CommonKVManager(BaseKVManager):
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
@@ -335,6 +342,28 @@ class CommonKVManager(BaseKVManager):
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(
@@ -1237,6 +1266,10 @@ class CommonKVSender(BaseKVSender):
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(
+94 -5
View File
@@ -1808,6 +1808,15 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
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)
@@ -2027,6 +2036,9 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
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
@@ -2064,11 +2076,23 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
)
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
@@ -2106,6 +2130,9 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
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
):
@@ -2125,9 +2152,67 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
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."""
@@ -2346,6 +2431,10 @@ class SchedulerDisaggregationDecodeMixin:
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)
+101 -27
View File
@@ -8,7 +8,7 @@ import struct
import threading
import time
from collections import defaultdict
from typing import List, Optional, Set, Tuple, Union
from typing import Dict, List, Optional, Set, Tuple, Union
import numpy as np
import numpy.typing as npt
@@ -214,6 +214,10 @@ class MooncakeKVManager(CommonKVManager):
# Per-room count of chunks not yet transferred; teardown waits for
# zero so a deferred chunk is not dropped by an early conclude.
self._staging_outstanding = defaultdict(int)
# Deferred KV release: aborted room -> (decode_ip, decode_port), ack
# held until the transfer drains. Written by the bootstrap thread,
# popped by the single transfer worker that owns the room.
self._deferred_ack_targets: Dict[int, Tuple[str, int]] = {}
self.session_lock = threading.Lock()
# Determine the number of threads to use for kv sender
cpu_count = os.cpu_count()
@@ -795,13 +799,7 @@ class MooncakeKVManager(CommonKVManager):
)
for (src_ptr, dst_ptr, item_len) in layers_params
]
for future in concurrent.futures.as_completed(futures):
status = future.result()
if status != 0:
for f in futures:
f.cancel()
return status
return 0
return self._await_transfer_futures(futures)
else:
# Combining all layers' params in one batch transfer is more efficient
# compared to using multiple threads
@@ -853,6 +851,25 @@ class MooncakeKVManager(CommonKVManager):
"must enable it and use the same page size and model spec."
)
def _await_transfer_futures(self, futures) -> int:
"""Await a chunk's per-layer RDMA writes; return the first non-zero status.
cancel() is a no-op for a running future, so with deferred release on we
still drain the running ones before returning (no write may outlive this
call, which the drain-ack relies on). Off: original early-return."""
ret = 0
for future in concurrent.futures.as_completed(futures):
try:
status = future.result()
except concurrent.futures.CancelledError:
continue
if status != 0 and ret == 0:
ret = status
for f in futures:
f.cancel()
if not self.enable_deferred_decode_kv_release:
return ret
return ret
def send_kvcache(
self,
mooncake_session_id: str,
@@ -978,13 +995,7 @@ class MooncakeKVManager(CommonKVManager):
executor.submit(process_layer, src_ptr, dst_ptr, token_item_len)
for src_ptr, dst_ptr, token_item_len in layers_params
]
for future in concurrent.futures.as_completed(futures):
status = future.result()
if status != 0:
for pending in futures:
pending.cancel()
return status
return 0
return self._await_transfer_futures(futures)
transfer_blocks = []
for src_ptr, dst_ptr, token_item_len in layers_params:
@@ -1112,14 +1123,7 @@ class MooncakeKVManager(CommonKVManager):
executor.submit(process_layer_tp_aware, src_v_ptrs[i], dst_v_ptrs[i])
)
for future in concurrent.futures.as_completed(futures):
status = future.result()
if status != 0:
for f in futures:
f.cancel()
return status
return 0
return self._await_transfer_futures(futures)
def send_aux(
self,
@@ -1612,6 +1616,39 @@ class MooncakeKVManager(CommonKVManager):
is_ipv6=na.is_ipv6,
)
def _prefill_unique_rank(self) -> int:
"""Stable per-sender id, matching what the transfer worker syncs on Success."""
return (
self.attn_tp_rank * (self.pp_size * self.attn_cp_size)
+ self.pp_rank * self.attn_cp_size
+ self.attn_cp_rank
)
def _send_abort_ack(self, decode_ip: str, decode_port: int, room: int) -> None:
"""Best-effort ack that this rank's transfer for an aborted room drained."""
try:
na = NetworkAddress(decode_ip, decode_port)
self._send_multipart_locked(
na.to_tcp(),
[
b"ABORT_ACK",
str(room).encode("ascii"),
str(self._prefill_unique_rank()).encode("ascii"),
],
is_ipv6=na.is_ipv6,
)
except Exception as e:
logger.debug(f"Failed to send drained ABORT_ACK for room {room}: {e}")
def _maybe_ack_drained_abort(self, room: int) -> None:
"""Send the deferred ack once an aborted room's chunks have drained
(outstanding == 0). pop() makes it fire at most once."""
if self._staging_outstanding.get(room, 0) > 0:
return
target = self._deferred_ack_targets.pop(room, None)
if target is not None:
self._send_abort_ack(target[0], target[1], room)
def transfer_worker(
self,
queue: FastQueue,
@@ -1651,6 +1688,9 @@ class MooncakeKVManager(CommonKVManager):
thread_finish_flag=True,
)
self._staging_outstanding.pop(kv_chunk.room, None)
if self.enable_deferred_decode_kv_release:
# Skipped => nothing written for this aborted room; ack.
self._maybe_ack_drained_abort(kv_chunk.room)
continue
# Count each chunk once; the flag survives re-enqueue on defer.
@@ -1925,6 +1965,10 @@ class MooncakeKVManager(CommonKVManager):
continue
self._staging_outstanding[kv_chunk.room] -= 1
if self.enable_deferred_decode_kv_release:
# In-flight write finished; if aborted and nothing outstanding,
# the pages are idle -> release the held ack.
self._maybe_ack_drained_abort(kv_chunk.room)
# Tear down only when no chunk is still outstanding and the room
# has concluded: already cleared, Success, or a Failed *last*
# chunk. A non-last Failed chunk keeps the room (more chunks may
@@ -1984,11 +2028,36 @@ class MooncakeKVManager(CommonKVManager):
room_to_be_aborted = int(waiting_req_bytes[1].decode("ascii"))
decode_ip = waiting_req_bytes[2].decode("ascii")
decode_port = int(waiting_req_bytes[3].decode("ascii"))
# No need to abort the room if it has already succeeded
if (
room_active = (
room_to_be_aborted in self.request_status
and self.check_status(room_to_be_aborted) != KVPoll.Success
):
)
if self.enable_deferred_decode_kv_release:
# Mark Failed FIRST (stops add_transfer_request enqueuing
# new chunks), THEN register the ack target: registering
# first would let the worker drain+ack while the room is
# not yet Failed, so a newly enqueued chunk could still
# write to the freed pages. The worker (not this thread)
# acks once its in-flight write drains; if nothing is in
# flight, decode falls back to the release timeout.
if room_active:
self.update_status(room_to_be_aborted, KVPoll.Failed)
self._deferred_ack_targets[room_to_be_aborted] = (
decode_ip,
decode_port,
)
logger.debug(
f"Received abort notification for room {room_to_be_aborted}, "
f"marked as Failed; ACK deferred until transfer drains"
)
else:
# Already completed/unknown: no in-flight write, ack now.
self._send_abort_ack(
decode_ip, decode_port, room_to_be_aborted
)
continue
# No need to abort the room if it has already succeeded
if room_active:
self.update_status(room_to_be_aborted, KVPoll.Failed)
logger.debug(
f"Received abort notification for room {room_to_be_aborted}, "
@@ -2102,9 +2171,14 @@ class MooncakeKVManager(CommonKVManager):
# Prefill acknowledges abort notification
if msg[0] == b"ABORT_ACK":
# TODO(shangming): use this info to implement the deferred release mechanism if needed
ack_aborted_room = int(msg[1].decode("ascii"))
logger.debug(f"Received ABORT_ACK for room {ack_aborted_room}")
# Deferred release: the 3-frame ack carries the prefill rank
# and means its transfer drained; aggregate for is_abort_release_safe.
if self.enable_deferred_decode_kv_release and len(msg) >= 3:
self.note_abort_ack(
ack_aborted_room, int(msg[2].decode("ascii"))
)
continue
bootstrap_room, status, prefill_rank = msg
+5
View File
@@ -619,6 +619,11 @@ class Envs:
SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False)
SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS = EnvInt(0)
SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL = EnvInt(120)
# Deferred decode-side KV release: on abort, hold an in-flight request's KV
# pages/slot until the prefill acks the transfer drained, or the timeout
# below fires. Off by default (no behavior/perf impact when disabled).
SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE = EnvBool(False)
SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE_TIMEOUT = EnvFloat(30.0)
# ===================================================================
# Distributed and model-parallel runtime
+23 -2
View File
@@ -4054,7 +4054,14 @@ class Scheduler(
# memory leak check (skipped for hisparse — pool counters intentionally
# diverge during host-backup, see _get_swa_token_info clamp).
if not self.enable_hisparse:
# Also skipped while deferred KV releases are pending: they hold pages out
# of the allocator by design, so the pool is transiently below `total` and
# would trip the idle leak invariant. Resumes once the holds resolve.
deferred_pending = (
self.disaggregation_mode == DisaggregationMode.DECODE
and self.disagg_decode_transfer_queue.has_pending_deferred_releases()
)
if not self.enable_hisparse and not deferred_pending:
has_leak, messages = self.invariant_checker._check_all_pools(
self.pool_stats_observer.get_pool_stats(),
)
@@ -4542,7 +4549,21 @@ class Scheduler(
for decode_req in self.disagg_decode_transfer_queue.queue:
if recv_req.abort_all or decode_req.req.rid.startswith(recv_req.rid):
logger.debug(f"Abort transfer queue request. {decode_req.req.rid=}")
decode_req.kv_receiver.abort()
receiver = decode_req.kv_receiver
receiver.abort()
# Arm drain-ack accounting once the ABORT is sent, so acks
# arriving before this req is deferred (e.g. during the next
# forward step) are captured. A fresh set also drops stale acks
# from a prior request that reused this bootstrap_room. A
# redundant abort only re-wipes -- holds longer, never releases
# early -- so no transition guard is needed.
if (
receiver.kv_mgr.enable_deferred_decode_kv_release
and receiver.abort_notified
):
receiver.kv_mgr.register_deferred_abort_room(
decode_req.req.bootstrap_room
)
# Abort requests whose KV is already backed up for retraction.
if self.disagg_decode_prealloc_queue.retracted_queue:
@@ -1468,6 +1468,9 @@ class SchedulerPPMixin:
def process_decode_transfer_queue(
self: Scheduler, release_rids: Optional[List[str]]
):
# Resolve held deferred releases every call, independent of release_rids,
# so ack/timeout-driven releases still fire when no rids are being polled.
self.disagg_decode_transfer_queue.resolve_deferred_releases()
if release_rids is not None:
released_reqs = self.disagg_decode_transfer_queue.pop_transferred(
release_rids