[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.is_hybrid_mla_backend = getattr(args, "is_hybrid_mla_backend", False)
self.disaggregation_mode = disaggregation_mode self.disaggregation_mode = disaggregation_mode
self.server_args = server_args 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 # for p/d multi node infer
self.bootstrap_host = server_args.host self.bootstrap_host = server_args.host
self.bootstrap_port = server_args.disaggregation_bootstrap_port self.bootstrap_port = server_args.disaggregation_bootstrap_port
@@ -251,6 +254,10 @@ class CommonKVManager(BaseKVManager):
self.session_pool_lock = threading.Lock() self.session_pool_lock = threading.Lock()
self.addr_to_rooms_tracker: Dict[str, Set[int]] = defaultdict(set) self.addr_to_rooms_tracker: Dict[str, Set[int]] = defaultdict(set)
self.prefill_response_tracker: Dict[int, 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 # Heartbeat interval should be at least 2 seconds
self.heartbeat_interval = max( self.heartbeat_interval = max(
envs.SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL.get(), 2.0 envs.SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL.get(), 2.0
@@ -335,6 +342,28 @@ class CommonKVManager(BaseKVManager):
with self.failure_lock: with self.failure_lock:
self.failure_records[bootstrap_room] = failure_reason 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: def get_kv_replica_factor(self) -> int:
if self._kv_replica_factor is None: if self._kv_replica_factor is None:
logger.warning_once( logger.warning_once(
@@ -1237,6 +1266,10 @@ class CommonKVSender(BaseKVSender):
self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, None) self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, None)
if hasattr(self.kv_mgr, "transfer_infos"): if hasattr(self.kv_mgr, "transfer_infos"):
self.kv_mgr.transfer_infos.pop(self.bootstrap_room, None) 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): def abort(self):
self.kv_mgr.record_failure( self.kv_mgr.record_failure(
+94 -5
View File
@@ -1808,6 +1808,15 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
self.spec_algorithm = scheduler.spec_algorithm self.spec_algorithm = scheduler.spec_algorithm
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
self.staging_handler = None 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: def add(self, decode_req: DecodeRequest) -> None:
self.queue.append(decode_req) self.queue.append(decode_req)
@@ -2027,6 +2036,9 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
transferred_reqs = [] transferred_reqs = []
indices_to_remove = set() 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)): 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: if rids_to_check is not None and decode_req.req.rid not in rids_to_check:
continue continue
@@ -2064,11 +2076,23 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
) )
if self.scheduler.enable_hisparse: if self.scheduler.enable_hisparse:
self.scheduler.hisparse_coordinator.request_finished(decode_req.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 if (
release_kv_cache(decode_req.req, self.tree_cache, is_insert=False) self.enable_deferred_kv_release
decode_req.kv_receiver.clear() and decode_req.kv_receiver.abort_notified
decode_req.kv_receiver = None ):
indices_to_remove.add(i) # 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: if self.scheduler.metrics_reporter.enable_metrics:
self.scheduler.metrics_collector.increment_transfer_failed_reqs() self.scheduler.metrics_collector.increment_transfer_failed_reqs()
continue continue
@@ -2106,6 +2130,9 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
raise ValueError(f"Unexpected poll case: {poll}") raise ValueError(f"Unexpected poll case: {poll}")
for i in indices_to_remove: 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( if self.enable_staging and self.staging_handler.is_staging_room(
self.queue[i].req.bootstrap_room self.queue[i].req.bootstrap_room
): ):
@@ -2125,9 +2152,67 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
return transferred_reqs 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): def release_memory_occupation(self):
"""Clean up in-flight transfers before releasing GPU memory.""" """Clean up in-flight transfers before releasing GPU memory."""
self.queue.clear() self.queue.clear()
# Pool is being torn down; drop held entries without per-request release.
self._deferred_releases.clear()
def resume_memory_occupation(self): def resume_memory_occupation(self):
"""Queues are already cleared on release; new transfers can be accepted.""" """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: if get_disagg().disaggregation_decode_enable_offload_kvcache:
self.decode_offload_manager.check_offload_progress() 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 # 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() resumed_reqs = self.disagg_decode_prealloc_queue.resume_retracted_reqs()
self.waiting_queue.extend(resumed_reqs) self.waiting_queue.extend(resumed_reqs)
+101 -27
View File
@@ -8,7 +8,7 @@ import struct
import threading import threading
import time import time
from collections import defaultdict 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 as np
import numpy.typing as npt import numpy.typing as npt
@@ -214,6 +214,10 @@ class MooncakeKVManager(CommonKVManager):
# Per-room count of chunks not yet transferred; teardown waits for # Per-room count of chunks not yet transferred; teardown waits for
# zero so a deferred chunk is not dropped by an early conclude. # zero so a deferred chunk is not dropped by an early conclude.
self._staging_outstanding = defaultdict(int) 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() self.session_lock = threading.Lock()
# Determine the number of threads to use for kv sender # Determine the number of threads to use for kv sender
cpu_count = os.cpu_count() cpu_count = os.cpu_count()
@@ -795,13 +799,7 @@ class MooncakeKVManager(CommonKVManager):
) )
for (src_ptr, dst_ptr, item_len) in layers_params for (src_ptr, dst_ptr, item_len) in layers_params
] ]
for future in concurrent.futures.as_completed(futures): return self._await_transfer_futures(futures)
status = future.result()
if status != 0:
for f in futures:
f.cancel()
return status
return 0
else: else:
# Combining all layers' params in one batch transfer is more efficient # Combining all layers' params in one batch transfer is more efficient
# compared to using multiple threads # compared to using multiple threads
@@ -853,6 +851,25 @@ class MooncakeKVManager(CommonKVManager):
"must enable it and use the same page size and model spec." "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( def send_kvcache(
self, self,
mooncake_session_id: str, mooncake_session_id: str,
@@ -978,13 +995,7 @@ class MooncakeKVManager(CommonKVManager):
executor.submit(process_layer, src_ptr, dst_ptr, token_item_len) executor.submit(process_layer, src_ptr, dst_ptr, token_item_len)
for src_ptr, dst_ptr, token_item_len in layers_params for src_ptr, dst_ptr, token_item_len in layers_params
] ]
for future in concurrent.futures.as_completed(futures): return self._await_transfer_futures(futures)
status = future.result()
if status != 0:
for pending in futures:
pending.cancel()
return status
return 0
transfer_blocks = [] transfer_blocks = []
for src_ptr, dst_ptr, token_item_len in layers_params: 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]) executor.submit(process_layer_tp_aware, src_v_ptrs[i], dst_v_ptrs[i])
) )
for future in concurrent.futures.as_completed(futures): return self._await_transfer_futures(futures)
status = future.result()
if status != 0:
for f in futures:
f.cancel()
return status
return 0
def send_aux( def send_aux(
self, self,
@@ -1612,6 +1616,39 @@ class MooncakeKVManager(CommonKVManager):
is_ipv6=na.is_ipv6, 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( def transfer_worker(
self, self,
queue: FastQueue, queue: FastQueue,
@@ -1651,6 +1688,9 @@ class MooncakeKVManager(CommonKVManager):
thread_finish_flag=True, thread_finish_flag=True,
) )
self._staging_outstanding.pop(kv_chunk.room, None) 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 continue
# Count each chunk once; the flag survives re-enqueue on defer. # Count each chunk once; the flag survives re-enqueue on defer.
@@ -1925,6 +1965,10 @@ class MooncakeKVManager(CommonKVManager):
continue continue
self._staging_outstanding[kv_chunk.room] -= 1 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 # Tear down only when no chunk is still outstanding and the room
# has concluded: already cleared, Success, or a Failed *last* # has concluded: already cleared, Success, or a Failed *last*
# chunk. A non-last Failed chunk keeps the room (more chunks may # 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")) room_to_be_aborted = int(waiting_req_bytes[1].decode("ascii"))
decode_ip = waiting_req_bytes[2].decode("ascii") decode_ip = waiting_req_bytes[2].decode("ascii")
decode_port = int(waiting_req_bytes[3].decode("ascii")) decode_port = int(waiting_req_bytes[3].decode("ascii"))
# No need to abort the room if it has already succeeded room_active = (
if (
room_to_be_aborted in self.request_status room_to_be_aborted in self.request_status
and self.check_status(room_to_be_aborted) != KVPoll.Success 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) self.update_status(room_to_be_aborted, KVPoll.Failed)
logger.debug( logger.debug(
f"Received abort notification for room {room_to_be_aborted}, " f"Received abort notification for room {room_to_be_aborted}, "
@@ -2102,9 +2171,14 @@ class MooncakeKVManager(CommonKVManager):
# Prefill acknowledges abort notification # Prefill acknowledges abort notification
if msg[0] == b"ABORT_ACK": 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")) ack_aborted_room = int(msg[1].decode("ascii"))
logger.debug(f"Received ABORT_ACK for room {ack_aborted_room}") 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 continue
bootstrap_room, status, prefill_rank = msg 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_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False)
SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS = EnvInt(0) SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS = EnvInt(0)
SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL = EnvInt(120) 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 # Distributed and model-parallel runtime
+23 -2
View File
@@ -4054,7 +4054,14 @@ class Scheduler(
# memory leak check (skipped for hisparse — pool counters intentionally # memory leak check (skipped for hisparse — pool counters intentionally
# diverge during host-backup, see _get_swa_token_info clamp). # 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( has_leak, messages = self.invariant_checker._check_all_pools(
self.pool_stats_observer.get_pool_stats(), self.pool_stats_observer.get_pool_stats(),
) )
@@ -4542,7 +4549,21 @@ class Scheduler(
for decode_req in self.disagg_decode_transfer_queue.queue: for decode_req in self.disagg_decode_transfer_queue.queue:
if recv_req.abort_all or decode_req.req.rid.startswith(recv_req.rid): 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=}") 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. # Abort requests whose KV is already backed up for retraction.
if self.disagg_decode_prealloc_queue.retracted_queue: if self.disagg_decode_prealloc_queue.retracted_queue:
@@ -1468,6 +1468,9 @@ class SchedulerPPMixin:
def process_decode_transfer_queue( def process_decode_transfer_queue(
self: Scheduler, release_rids: Optional[List[str]] 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: if release_rids is not None:
released_reqs = self.disagg_decode_transfer_queue.pop_transferred( released_reqs = self.disagg_decode_transfer_queue.pop_transferred(
release_rids release_rids
@@ -211,6 +211,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
queue = DecodeTransferQueue.__new__(DecodeTransferQueue) queue = DecodeTransferQueue.__new__(DecodeTransferQueue)
queue.queue = [decode_req] queue.queue = [decode_req]
queue.enable_staging = False queue.enable_staging = False
queue.enable_deferred_kv_release = False
queue.gloo_group = MagicMock() queue.gloo_group = MagicMock()
queue.req_to_metadata_buffer_idx_allocator = MagicMock() queue.req_to_metadata_buffer_idx_allocator = MagicMock()
queue.tp_rank = 0 queue.tp_rank = 0
@@ -0,0 +1,247 @@
"""Unit tests for the deferred decode-side KV release mechanism.
When a decode request is aborted while its prefill->decode KV transfer may still
be in flight, the decode side holds its KV pages / req-slot instead of freeing
them immediately (which could let the still-in-flight write land on pages already
reused by another request). The pages are released once every prefill rank acks
that its transfer drained (CommonKVManager.is_abort_release_safe), or a timeout
fires. See DecodeTransferQueue.resolve_deferred_releases.
"""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.disaggregation import decode as decode_mod
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.decode import DecodeTransferQueue
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _make_manager():
"""A bare CommonKVManager carrying only the deferred-ack state the helpers
touch (avoids the heavy real __init__)."""
mgr = CommonKVManager.__new__(CommonKVManager)
mgr._deferred_abort_ack_tracker = {}
return mgr
class TestAbortAckAggregation(CustomTestCase):
def test_release_safe_only_after_all_required_ranks_ack(self):
mgr = _make_manager()
room = 100
mgr.register_deferred_abort_room(room)
self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2))
mgr.note_abort_ack(room, 0)
self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2))
mgr.note_abort_ack(room, 1)
self.assertTrue(mgr.is_abort_release_safe(room, required_acks=2))
def test_duplicate_rank_ack_does_not_over_count(self):
mgr = _make_manager()
room = 101
mgr.register_deferred_abort_room(room)
mgr.note_abort_ack(room, 0)
mgr.note_abort_ack(room, 0) # same rank twice
# Two acks arrived but from one rank: not safe for a 2-rank prefill.
self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2))
def test_single_rank_fast_path(self):
mgr = _make_manager()
room = 102
mgr.register_deferred_abort_room(room)
mgr.note_abort_ack(room, 0)
self.assertTrue(mgr.is_abort_release_safe(room, required_acks=1))
def test_clear_deferred_abort_state(self):
mgr = _make_manager()
room = 103
mgr.register_deferred_abort_room(room)
mgr.note_abort_ack(room, 0)
mgr.clear_deferred_abort_state(room)
self.assertNotIn(room, mgr._deferred_abort_ack_tracker)
self.assertFalse(mgr.is_abort_release_safe(room, required_acks=1))
def test_ack_before_register_is_dropped(self):
# An ack for a room that isn't actively held must not be recorded (it
# would otherwise pollute a later request reusing the same room).
mgr = _make_manager()
room = 104
mgr.note_abort_ack(room, 0) # no register yet
self.assertNotIn(room, mgr._deferred_abort_ack_tracker)
self.assertFalse(mgr.is_abort_release_safe(room, required_acks=1))
def test_late_ack_after_release_does_not_pollute_reused_room(self):
# Regression for bootstrap_room reuse: req A (room R) releases, then a
# late ack from A arrives, then req B reuses room R. B must start from a
# clean slate and not inherit A's ack (which would release B early while
# its transfer is still in flight -> KV corruption).
mgr = _make_manager()
room = 105
# Req A: held, one of two ranks acks, then released (e.g. timed out).
mgr.register_deferred_abort_room(room)
mgr.note_abort_ack(room, 0)
mgr.clear_deferred_abort_state(room)
# Late ack from A's other rank arrives after release -> dropped.
mgr.note_abort_ack(room, 1)
self.assertNotIn(room, mgr._deferred_abort_ack_tracker)
# Req B reuses room R.
mgr.register_deferred_abort_room(room)
# Only B's rank-0 has acked so far; a 2-rank prefill is NOT safe yet.
mgr.note_abort_ack(room, 0)
self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2))
mgr.note_abort_ack(room, 1)
self.assertTrue(mgr.is_abort_release_safe(room, required_acks=2))
def test_register_resets_stale_acks(self):
mgr = _make_manager()
room = 106
mgr.register_deferred_abort_room(room)
mgr.note_abort_ack(room, 0)
mgr.note_abort_ack(room, 1)
self.assertTrue(mgr.is_abort_release_safe(room, required_acks=2))
# Re-registering (a later reuse) wipes the prior acks.
mgr.register_deferred_abort_room(room)
self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2))
class _FakeIdxAllocator:
def __init__(self):
self.freed = []
def free(self, idx):
self.freed.append(idx)
def _make_queue(timeout=30.0):
q = DecodeTransferQueue.__new__(DecodeTransferQueue)
q._deferred_releases = []
q.deferred_kv_release_timeout = timeout
q.enable_staging = False
q.staging_handler = None
q.tree_cache = object()
q.metadata_buffers = SimpleNamespace(bootstrap_room={})
q.req_to_metadata_buffer_idx_allocator = _FakeIdxAllocator()
return q
def _make_decode_req(room, idx, mgr, n_prefill_ranks=1):
receiver = SimpleNamespace(
kv_mgr=mgr,
# One entry per prefill rank the decode notified of the abort; its length
# is the required drain-ack count (see DecodeTransferQueue._defer_release).
bootstrap_infos=[{"rank": r} for r in range(n_prefill_ranks)],
clear=lambda: None,
)
return SimpleNamespace(
req=SimpleNamespace(bootstrap_room=room),
kv_receiver=receiver,
metadata_buffer_index=idx,
)
class TestResolveDeferredReleases(CustomTestCase):
def test_noop_when_nothing_deferred(self):
q = _make_queue()
with patch.object(decode_mod, "release_kv_cache") as rel:
q.resolve_deferred_releases()
rel.assert_not_called()
def test_holds_until_drained_then_releases(self):
mgr = _make_manager()
room, idx = 200, 7
q = _make_queue()
dreq = _make_decode_req(room, idx, mgr, n_prefill_ranks=2)
# In production the room is armed in abort_request when the ABORT is
# sent, before the scheduler defers here.
mgr.register_deferred_abort_room(room)
q._defer_release(dreq)
with patch.object(decode_mod, "release_kv_cache") as rel:
# Not yet acked -> held, not released.
q.resolve_deferred_releases()
rel.assert_not_called()
self.assertEqual(len(q._deferred_releases), 1)
# One of two ranks acked -> still held.
mgr.note_abort_ack(room, 0)
q.resolve_deferred_releases()
rel.assert_not_called()
self.assertEqual(len(q._deferred_releases), 1)
# Both ranks acked -> released exactly once.
mgr.note_abort_ack(room, 1)
q.resolve_deferred_releases()
rel.assert_called_once_with(dreq.req, q.tree_cache, is_insert=False)
# Held state fully cleaned up.
self.assertEqual(q._deferred_releases, [])
self.assertEqual(q.req_to_metadata_buffer_idx_allocator.freed, [idx])
self.assertEqual(q.metadata_buffers.bootstrap_room[idx], 0)
self.assertNotIn(room, mgr._deferred_abort_ack_tracker)
self.assertIsNone(dreq.kv_receiver)
def test_releases_on_timeout_without_ack(self):
mgr = _make_manager()
room, idx = 300, 3
q = _make_queue(timeout=30.0)
dreq = _make_decode_req(room, idx, mgr, n_prefill_ranks=1)
# Force an already-expired deadline (no ack will ever arrive).
q._deferred_releases.append((dreq, float("-inf"), idx, 1))
with patch.object(decode_mod, "release_kv_cache") as rel:
q.resolve_deferred_releases()
rel.assert_called_once_with(dreq.req, q.tree_cache, is_insert=False)
self.assertEqual(q._deferred_releases, [])
self.assertEqual(q.req_to_metadata_buffer_idx_allocator.freed, [idx])
self.assertIsNone(dreq.kv_receiver)
def test_failed_release_is_isolated_and_not_retried(self):
# A raising _do_release must drop the entry (no double-free on retry) and
# not brick resolve for the remaining entries or subsequent calls.
mgr = _make_manager()
q = _make_queue()
good = _make_decode_req(700, 1, mgr)
bad = _make_decode_req(701, 2, mgr)
# Both already past deadline -> both selected for release.
q._deferred_releases.append((bad, float("-inf"), 2, 1))
q._deferred_releases.append((good, float("-inf"), 1, 1))
calls = []
def fake_release(req, tree_cache, is_insert):
calls.append(req)
if req is bad.req:
raise RuntimeError("boom")
with patch.object(decode_mod, "release_kv_cache", side_effect=fake_release):
q.resolve_deferred_releases() # must not raise
# The good one still released despite the bad one throwing.
self.assertIn(good.req, calls)
# Nothing left held, and a second call is a clean no-op (no retry).
self.assertEqual(q._deferred_releases, [])
q.resolve_deferred_releases()
def test_defer_release_records_deadline_and_idx(self):
mgr = _make_manager()
q = _make_queue(timeout=12.5)
dreq = _make_decode_req(room=400, idx=9, mgr=mgr)
q._defer_release(dreq)
self.assertEqual(len(q._deferred_releases), 1)
held_req, deadline, held_idx, required = q._deferred_releases[0]
self.assertIs(held_req, dreq)
self.assertEqual(held_idx, 9)
self.assertIsInstance(deadline, float)
if __name__ == "__main__":
unittest.main()