From de816e1eb53e12e1538df47d300ad0a74e63c54e Mon Sep 17 00:00:00 2001
From: YAMY <74099316+YAMY1234@users.noreply.github.com>
Date: Fri, 24 Jul 2026 02:22:55 -0700
Subject: [PATCH] [Disagg][StagingBuffer][1/2] Robustness and failure handling
(#31217)
---
.../advanced_features/pd_disaggregation.mdx | 6 --
.../docs/references/environment_variables.mdx | 5 --
.../disaggregation/common/staging_buffer.py | 3 +
.../disaggregation/common/staging_handler.py | 87 +++++++++++++-----
python/sglang/srt/disaggregation/decode.py | 25 +++---
.../srt/disaggregation/mooncake/conn.py | 89 ++++++++++---------
python/sglang/srt/disaggregation/nixl/conn.py | 48 ++++++----
python/sglang/srt/disaggregation/prefill.py | 33 +++++--
python/sglang/srt/environ.py | 1 -
.../test_disaggregation_different_tp.py | 3 +-
.../disaggregation/test_nixl_backend_basic.py | 1 +
11 files changed, 187 insertions(+), 114 deletions(-)
diff --git a/docs_new/docs/advanced_features/pd_disaggregation.mdx b/docs_new/docs/advanced_features/pd_disaggregation.mdx
index 0c604d7eb..d2a781042 100644
--- a/docs_new/docs/advanced_features/pd_disaggregation.mdx
+++ b/docs_new/docs/advanced_features/pd_disaggregation.mdx
@@ -261,11 +261,6 @@ Enable the staging buffer when prefill and decode use **different TP sizes** wit
Enable GPU staging buffer for heterogeneous TP KV transfer |
False |
-
- SGLANG_DISAGG_STAGING_BUFFER_SIZE_MB |
- Prefill-side per-worker staging buffer size in MB |
- 64 |
-
SGLANG_DISAGG_STAGING_POOL_SIZE_MB |
Decode-side ring buffer pool total size in MB |
@@ -279,7 +274,6 @@ Enable the staging buffer when prefill and decode use **different TP sizes** wit
```bash Command
# Set staging buffer environment variables on BOTH prefill and decode
export SGLANG_DISAGG_STAGING_BUFFER=1
-export SGLANG_DISAGG_STAGING_BUFFER_SIZE_MB=64
export SGLANG_DISAGG_STAGING_POOL_SIZE_MB=4096
# Prefill with TP=4
diff --git a/docs_new/docs/references/environment_variables.mdx b/docs_new/docs/references/environment_variables.mdx
index 89870857f..db244e8e1 100644
--- a/docs_new/docs/references/environment_variables.mdx
+++ b/docs_new/docs/references/environment_variables.mdx
@@ -908,11 +908,6 @@ SGLang supports various environment variables that can be used to configure its
Enable GPU staging buffer for heterogeneous TP KV transfer. Required when prefill and decode use different TP/attention-TP sizes. Only for non-MLA models (e.g. GQA, MHA). |
false |
-
- SGLANG_DISAGG_STAGING_BUFFER_SIZE_MB |
- Prefill-side per-worker staging buffer size in MB. Used for gathering KV head slices before bulk RDMA transfer. |
- 64 |
-
SGLANG_DISAGG_STAGING_POOL_SIZE_MB |
Decode-side ring buffer pool total size in MB. Shared buffer receiving RDMA data from all prefill ranks. Larger values support higher concurrency. |
diff --git a/python/sglang/srt/disaggregation/common/staging_buffer.py b/python/sglang/srt/disaggregation/common/staging_buffer.py
index 19eccbf4a..e41071cba 100644
--- a/python/sglang/srt/disaggregation/common/staging_buffer.py
+++ b/python/sglang/srt/disaggregation/common/staging_buffer.py
@@ -194,6 +194,9 @@ class StagingAllocator:
self.watermark_round = 0
self.watermark_tail = 0
self.lock = threading.Lock()
+ # Lazily created on the decode side by the first scatter; stays None
+ # until then so release_room can drain it without a defensive check.
+ self._scatter_stream = None
logger.info(
f"StagingAllocator (ring+overcommit): "
diff --git a/python/sglang/srt/disaggregation/common/staging_handler.py b/python/sglang/srt/disaggregation/common/staging_handler.py
index a0b6bcbfc..b2ef46ad1 100644
--- a/python/sglang/srt/disaggregation/common/staging_handler.py
+++ b/python/sglang/srt/disaggregation/common/staging_handler.py
@@ -18,6 +18,10 @@ import torch
logger = logging.getLogger(__name__)
+# Bounded wait for a watermark advance before re-enqueueing a deferred staging
+# chunk, so the re-enqueue retry does not busy-spin a core.
+STAGING_WATERMARK_WAIT_S = 0.001
+
if TYPE_CHECKING:
from sglang.srt.disaggregation.decode import DecodeRequest
@@ -79,11 +83,14 @@ class DecodeStagingHandler:
self.tp_rank = tp_rank
self.scheduler = scheduler
self._room_to_decode_req: dict = {}
+ # Stashed at registration: removal paths null decode_req.kv_receiver
+ # before unregister runs, but release_room still needs it.
+ self._room_to_receiver: dict = {}
self._wm_subscribers: dict = {}
def register_wm_subscriber(self, receiver, session_id: str) -> None:
"""Register a prefill's bootstrap connection for watermark broadcasts."""
- if receiver is None or not getattr(receiver, "bootstrap_infos", None):
+ if receiver is None or not receiver.bootstrap_infos:
return
key = tuple(str(bi) for bi in receiver.bootstrap_infos)
if key not in self._wm_subscribers:
@@ -133,10 +140,49 @@ class DecodeStagingHandler:
# ------------------------------------------------------------------
def register_decode_req(self, room: int, decode_req: DecodeRequest) -> None:
+ # Called once per room from pop_preallocated, before send_metadata.
+ decode_req._staging_scatter_done = False
+ decode_req._chunk_events = []
self._room_to_decode_req[room] = decode_req
+ self._room_to_receiver[room] = decode_req.kv_receiver
def unregister_decode_req(self, room: int) -> None:
- self._room_to_decode_req.pop(room, None)
+ # Pop before release_room so no new arrival can start consuming the slots.
+ decode_req = self._room_to_decode_req.pop(room, None)
+ receiver = self._room_to_receiver.pop(room, None)
+ if decode_req is not None:
+ self.release_room(room, decode_req, receiver)
+ self.kv_manager._staging_ctx.room_receivers.pop(room, None)
+ self.kv_manager._staging_ctx.room_bootstrap.pop(room, None)
+
+ def release_room(self, room: int, decode_req: DecodeRequest, receiver) -> None:
+ """Free outstanding staging allocations of a room; no-op after a
+ clean Success, releases watermark-pinning leaks on failure/abort."""
+ # Drain in-flight scatters before freeing anything, including one whose
+ # event is not yet in _chunk_events (submit_chunk_scatter records it
+ # after launching the kernel), so no scatter reads a freed staging slot
+ # or writes into KV-pool pages the failure path frees for reuse.
+ stream = self.staging_allocator._scatter_stream
+ if stream is not None:
+ stream.synchronize()
+ chunk_infos = receiver.chunk_staging_infos if receiver is not None else []
+ unscattered_allocs = []
+ for chunk_idx, info in enumerate(chunk_infos):
+ if info[0] >= 0:
+ unscattered_allocs.append((chunk_idx, info[0]))
+ chunk_infos[chunk_idx] = (-1, -1, 0, -1, 0)
+ for chunk_idx, alloc_id in unscattered_allocs:
+ logger.warning(
+ "[STAGING] releasing unscattered staging allocation "
+ "room=%s chunk=%s alloc_id=%s",
+ room,
+ chunk_idx,
+ alloc_id,
+ )
+ self._free_and_send_watermark(alloc_id, decode_req)
+ for _event, alloc_id in decode_req._chunk_events:
+ self._free_and_send_watermark(alloc_id, decode_req)
+ decode_req._chunk_events.clear()
# ------------------------------------------------------------------
# Scatter submission: called from decode_thread (background)
@@ -160,7 +206,8 @@ class DecodeStagingHandler:
chunk_idx,
)
return False
- chunk_infos = getattr(decode_req.kv_receiver, "chunk_staging_infos", [])
+ receiver = self._room_to_receiver.get(room)
+ chunk_infos = receiver.chunk_staging_infos if receiver is not None else []
if chunk_idx >= len(chunk_infos):
return False
alloc_id, staging_offset, _, _, _ = chunk_infos[chunk_idx]
@@ -171,8 +218,8 @@ class DecodeStagingHandler:
if ok:
event = torch.cuda.Event()
event.record(self.staging_allocator._scatter_stream)
- if not hasattr(decode_req, "_chunk_events"):
- decode_req._chunk_events = []
+ # Append before zeroing so the completion check always sees either
+ # the slot or the event.
decode_req._chunk_events.append((event, alloc_id))
chunk_infos[chunk_idx] = (-1, -1, 0, -1, 0)
else:
@@ -253,9 +300,7 @@ class DecodeStagingHandler:
def is_done(self, decode_req: DecodeRequest) -> bool:
"""Return True if staging scatter is complete for this request."""
- if not getattr(decode_req, "_staging_scatter_done", False):
- return False
- return not getattr(decode_req, "_chunk_events", None)
+ return decode_req._staging_scatter_done and not decode_req._chunk_events
def advance_scatter(self, decode_req: DecodeRequest) -> None:
"""Check CUDA events and free completed staging allocations.
@@ -312,7 +357,7 @@ class DecodeStagingHandler:
device = k_buffers[0].device
torch.cuda.set_device(device)
- if not hasattr(self.staging_allocator, "_scatter_stream"):
+ if self.staging_allocator._scatter_stream is None:
self.staging_allocator._scatter_stream = torch.cuda.Stream(device=device)
scatter_stream = self.staging_allocator._scatter_stream
@@ -350,7 +395,7 @@ class DecodeStagingHandler:
def _submit_last_scatter(self, decode_req: DecodeRequest) -> int:
"""Submit scatter for the last chunk. Returns alloc_id >= 0, or -1."""
receiver = decode_req.kv_receiver
- chunk_infos = getattr(receiver, "chunk_staging_infos", [])
+ chunk_infos = receiver.chunk_staging_infos if receiver is not None else []
if not chunk_infos:
return -1
@@ -601,22 +646,18 @@ def _get_custom_mem_pool(device: str):
return custom_mem_pool, pool_type
-def init_staging_buffers(register_fn, kv_args, count: int) -> list:
- """Create prefill-side staging buffers and register them with the transport.
+def init_staging_buffers(
+ register_fn, kv_args, count: int, chunked_prefill_size: int
+) -> list:
+ """Create prefill-side staging buffers, each sized to one prefill chunk.
- Args:
- register_fn: callable(ptr: int, size: int) that registers a memory
- region with the transport backend.
- kv_args: KVArgs with gpu_id.
- count: number of staging buffers to create.
-
- Returns list of StagingBuffer instances.
+ Sizing to one chunk (``chunked_prefill_size`` tokens of this rank's KV) means
+ a chunk can never be too large for the buffer.
"""
from sglang.srt.disaggregation.common.staging_buffer import StagingBuffer
- from sglang.srt.environ import envs
- size_mb = envs.SGLANG_DISAGG_STAGING_BUFFER_SIZE_MB.get()
- size_bytes = size_mb * 1024 * 1024
+ full_chunk_pages = max(1, chunked_prefill_size // kv_args.page_size)
+ size_bytes = full_chunk_pages * sum(kv_args.kv_item_lens)
gpu_id = kv_args.gpu_id
device = f"cuda:{gpu_id}"
@@ -693,7 +734,7 @@ def handle_staging_req(
session_id,
)
return
- infos = getattr(receiver, "chunk_staging_infos", [])
+ infos = receiver.chunk_staging_infos
if chunk_idx < len(infos) and infos[chunk_idx][0] >= 0:
_, offset, rnd, end, _ = infos[chunk_idx]
diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py
index 5224920c4..676165ba3 100644
--- a/python/sglang/srt/disaggregation/decode.py
+++ b/python/sglang/srt/disaggregation/decode.py
@@ -1172,6 +1172,16 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
page_indices = kv_to_page_indices(kv_indices, kv_transfer_page_size).astype(
np.int32
)
+ if (
+ self.transfer_queue.enable_staging
+ and hasattr(decode_req.kv_receiver, "require_staging")
+ and decode_req.kv_receiver.require_staging
+ ):
+ # Register before send_metadata, which triggers the STAGING_REQ
+ # prefetch (dropped for an unregistered room); tiny race, correct order.
+ self.transfer_queue.staging_handler.register_decode_req(
+ decode_req.req.bootstrap_room, decode_req
+ )
decode_req.kv_receiver.send_metadata(
page_indices,
decode_req.metadata_buffer_index,
@@ -1183,14 +1193,6 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
decode_req.kv_receiver,
decode_req.req.build_rebootstrap_payload(),
)
- if (
- self.transfer_queue.enable_staging
- and hasattr(decode_req.kv_receiver, "require_staging")
- and decode_req.kv_receiver.require_staging
- ):
- self.transfer_queue.staging_handler.register_decode_req(
- decode_req.req.bootstrap_room, decode_req
- )
preallocated_reqs.append(decode_req)
indices_to_remove.add(i)
decode_req.req.time_stats.set_decode_transfer_queue_entry_time()
@@ -1659,13 +1661,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
def extend(self, decode_reqs: List[DecodeRequest]) -> None:
self.queue.extend(decode_reqs)
- if self.enable_staging:
- for dr in decode_reqs:
- if (
- hasattr(dr.kv_receiver, "require_staging")
- and dr.kv_receiver.require_staging
- ):
- self.staging_handler.register_decode_req(dr.req.bootstrap_room, dr)
def _commit_transfer_to_req(self, decode_req: DecodeRequest):
idx = decode_req.metadata_buffer_index
diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py
index 7d306a4d0..fcd4db866 100644
--- a/python/sglang/srt/disaggregation/mooncake/conn.py
+++ b/python/sglang/srt/disaggregation/mooncake/conn.py
@@ -23,6 +23,7 @@ from sglang.srt.disaggregation.common.conn import (
KVTransferError,
)
from sglang.srt.disaggregation.common.staging_handler import (
+ STAGING_WATERMARK_WAIT_S,
DecodeStagingContext,
PrefillStagingContext,
StagingRegisterInfo,
@@ -299,6 +300,7 @@ class MooncakeKVManager(CommonKVManager):
lambda ptr, size: self.engine.batch_register([ptr], [size]),
self.kv_args,
count,
+ self.server_args.chunked_prefill_size,
)
self.kv_buffer_tensors = None
@@ -367,24 +369,21 @@ class MooncakeKVManager(CommonKVManager):
def _send_chunk_ready(self, req, chunk_idx, kv_chunk, prefill_unique_rank):
"""Notify decode that a non-last staging chunk RDMA is complete."""
- try:
- na = NetworkAddress(req.endpoint, req.dst_port)
- self._connect(
- na.to_tcp(),
- is_ipv6=na.is_ipv6,
- ).send_multipart(
- [
- b"CHUNK_READY",
- str(req.room).encode("ascii"),
- str(chunk_idx).encode("ascii"),
- str(kv_chunk.index_slice.start).encode("ascii"),
- str(len(kv_chunk.prefill_kv_indices)).encode("ascii"),
- req.mooncake_session_id.encode("ascii"),
- str(prefill_unique_rank).encode("ascii"),
- ]
- )
- except Exception:
- pass
+ na = NetworkAddress(req.endpoint, req.dst_port)
+ self._connect(
+ na.to_tcp(),
+ is_ipv6=na.is_ipv6,
+ ).send_multipart(
+ [
+ b"CHUNK_READY",
+ str(req.room).encode("ascii"),
+ str(chunk_idx).encode("ascii"),
+ str(kv_chunk.index_slice.start).encode("ascii"),
+ str(len(kv_chunk.prefill_kv_indices)).encode("ascii"),
+ req.mooncake_session_id.encode("ascii"),
+ str(prefill_unique_rank).encode("ascii"),
+ ]
+ )
def _do_staging_transfer(
self,
@@ -399,10 +398,11 @@ class MooncakeKVManager(CommonKVManager):
):
"""Execute staging transfer for one chunk. Returns (ret, deferred).
- Handles readiness check, transfer, fallback, and CHUNK_READY notification.
- deferred=True means caller should re-enqueue and break.
+ Handles readiness check, transfer, and CHUNK_READY notification; a chunk
+ that cannot fit returns -1 (the caller fails only this room) instead of
+ falling back to the slice path, which would leak the decode-side
+ allocation. deferred=True means caller should re-enqueue and break.
"""
- _tp = self.attn_tp_rank
ready, chunk_idx, c_offset, _, _ = staging_strategy.check_ready(
req,
kv_chunk.index_slice.start,
@@ -412,11 +412,19 @@ class MooncakeKVManager(CommonKVManager):
from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator
if c_offset == StagingAllocator.ALLOC_OVERSIZED:
- raise RuntimeError(
- f"[Staging] Chunk staging allocation permanently failed: "
- f"chunk exceeds ring buffer total size (room={kv_chunk.room}). "
- f"Increase SGLANG_DISAGG_STAGING_POOL_SIZE_MB."
+ # Fail this room, not the worker thread: the same prefill still
+ # serves other (same-TP, non-staging) decode instances.
+ logger.warning_once(
+ "[Staging] a chunk exceeds the staging ring; failing affected "
+ "requests. Increase SGLANG_DISAGG_STAGING_POOL_SIZE_MB or "
+ "reduce chunked_prefill_size."
)
+ return (-1, False)
+ # Not ready yet: wait (bounded) for a watermark advance, then
+ # re-enqueue to retry. A plain block-until-ready would head-of-line
+ # block other rooms on this single worker thread.
+ with self._staging_ctx.watermark_cv:
+ self._staging_ctx.watermark_cv.wait(STAGING_WATERMARK_WAIT_S)
queue.put(kv_chunk)
return (-1, True)
@@ -428,21 +436,15 @@ class MooncakeKVManager(CommonKVManager):
target_info,
)
if ret == -1:
- logger.warning(
- f"[Staging][tp{_tp}] Falling back to per-token slice path "
- f"(room={kv_chunk.room})"
+ # Doesn't fit the ring: fail this room (caller's ret != 0 path), do
+ # not fall back to the slice path (leaks the decode-side allocation).
+ logger.warning_once(
+ "[Staging] a chunk does not fit the staging ring; failing affected "
+ "requests. Increase SGLANG_DISAGG_STAGING_POOL_SIZE_MB or "
+ "reduce chunked_prefill_size."
)
- ret = self.send_kvcache_slice(
- req.mooncake_session_id,
- kv_chunk.prefill_kv_indices,
- target_info.dst_kv_ptrs,
- chunked_dst_kv_indice,
- target_info.dst_tp_rank,
- target_info.dst_attn_tp_size,
- target_info.dst_kv_item_len,
- executor,
- )
- elif ret == 0 and not kv_chunk.is_last_chunk:
+ return (-1, False)
+ if ret == 0 and not kv_chunk.is_last_chunk:
self._send_chunk_ready(req, chunk_idx, kv_chunk, prefill_unique_rank)
return (ret, False)
@@ -1493,6 +1495,13 @@ class MooncakeKVManager(CommonKVManager):
if kv_chunk.room in self.transfer_infos:
self.transfer_infos.pop(kv_chunk.room)
self.req_to_decode_prefix_len.pop(kv_chunk.room, None)
+ if self.enable_staging:
+ # Purge prefetch bookkeeping for the finished room.
+ # Snapshot first: the scheduler thread adds concurrently.
+ for key in list(self._staging_ctx.prefetch_requested):
+ if key[0] == kv_chunk.room:
+ self._staging_ctx.prefetch_requested.discard(key)
+ self._staging_ctx.prefetched_rooms.discard(kv_chunk.room)
except Exception as e:
# NOTE(shangming): Remove this when we make sure the transfer thread is bug-free
@@ -1957,11 +1966,11 @@ class MooncakeKVReceiver(CommonKVReceiver):
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return
+ self.chunk_staging_infos = []
if (
self.kv_mgr.enable_staging
and self.kv_mgr._staging_ctx.allocator is not None
):
- self.chunk_staging_infos = []
self.kv_mgr.register_staging_room_bootstrap(
self.bootstrap_room, self.bootstrap_infos, self
)
diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py
index f60a8766a..c7f4e4824 100644
--- a/python/sglang/srt/disaggregation/nixl/conn.py
+++ b/python/sglang/srt/disaggregation/nixl/conn.py
@@ -24,7 +24,10 @@ from sglang.srt.disaggregation.common.conn import (
CommonKVSender,
KVTransferError,
)
-from sglang.srt.disaggregation.common.staging_handler import StagingRegisterInfo
+from sglang.srt.disaggregation.common.staging_handler import (
+ STAGING_WATERMARK_WAIT_S,
+ StagingRegisterInfo,
+)
from sglang.srt.disaggregation.common.utils import (
FastQueue,
TransferKVChunk,
@@ -494,6 +497,7 @@ class NixlKVManager(CommonKVManager):
lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id),
self.kv_args,
count,
+ self.server_args.chunked_prefill_size,
)
def _init_staging_allocator(self):
@@ -1102,10 +1106,6 @@ class NixlKVManager(CommonKVManager):
# pick it up again on the next pop.
staging_deferred = True
break
- # kv_xfer_handle is None here means staging
- # send_kvcache_staged() returned None (e.g.
- # decode buffer too small) -- fall through to
- # the slice path below.
if kv_xfer_handle is None:
if self.is_mla_backend or (
@@ -1212,11 +1212,10 @@ class NixlKVManager(CommonKVManager):
self.req_to_decode_prefix_len.pop(room, None)
if self.enable_staging and self._staging_ctx is not None:
self._staging_ctx.prefetched_rooms.discard(room)
- self._staging_ctx.prefetch_requested = {
- k
- for k in self._staging_ctx.prefetch_requested
- if k[0] != room
- }
+ # Snapshot first: the scheduler thread adds concurrently.
+ for k in list(self._staging_ctx.prefetch_requested):
+ if k[0] == room:
+ self._staging_ctx.prefetch_requested.discard(k)
else:
self.update_status(room, KVPoll.Transferring)
except Exception as e:
@@ -1728,9 +1727,9 @@ class NixlKVManager(CommonKVManager):
- staging successfully posted -> return ``(handle, False)``. The
caller appends the handle to the per-chunk handle list and
busy-polls it to DONE alongside other handles.
- - send_kvcache_staged returned None (decode buffer too small,
- kv_buffer_tensors missing, etc.) -> return ``(None, False)``,
- signalling the caller to fall back to send_kvcache_slice.
+ - send_kvcache_staged returned None (chunk cannot fit; decode buffer
+ too small, kv_buffer_tensors missing, etc.) -> raise RuntimeError
+ instead of falling back to the slice path.
"""
page_start = kv_chunk.index_slice.start
num_pages = len(kv_chunk.prefill_kv_indices)
@@ -1750,6 +1749,11 @@ class NixlKVManager(CommonKVManager):
f"(room={kv_chunk.room}). Increase "
f"SGLANG_DISAGG_STAGING_POOL_SIZE_MB."
)
+ # Not ready yet: wait (bounded) for a watermark advance, then
+ # re-enqueue to retry. A plain block-until-ready would head-of-line
+ # block other rooms on this single worker thread.
+ with self._staging_ctx.watermark_cv:
+ self._staging_ctx.watermark_cv.wait(STAGING_WATERMARK_WAIT_S)
queue.put(kv_chunk)
return (None, True)
@@ -1770,6 +1774,17 @@ class NixlKVManager(CommonKVManager):
notif_tag,
staging_buffer=staging_strategy.staging_buffer,
)
+ if handle is None:
+ # A silent slice fallback would leak this chunk's decode-side
+ # allocation and pin the ring watermark; with grid-aligned sends
+ # not fitting can only mean misconfiguration.
+ raise RuntimeError(
+ f"[Staging] Staged transfer cannot fit chunk "
+ f"(room={kv_chunk.room}, chunk_idx={chunk_idx}, "
+ f"pages={num_pages}). Increase "
+ f"SGLANG_DISAGG_STAGING_POOL_SIZE_MB or reduce "
+ f"chunked_prefill_size."
+ )
return (handle, False)
def send_aux(
@@ -2525,10 +2540,7 @@ class NixlKVSender(CommonKVSender):
def clear(self) -> None:
super().clear()
- if (
- getattr(self.kv_mgr, "enable_staging", False)
- and getattr(self.kv_mgr, "_staging_ctx", None) is not None
- ):
+ if self.kv_mgr.enable_staging and self.kv_mgr._staging_ctx is not None:
self.kv_mgr._staging_ctx.prefetched_rooms.discard(self.bootstrap_room)
self.kv_mgr._staging_ctx.prefetch_requested = {
key
@@ -2584,11 +2596,11 @@ class NixlKVReceiver(CommonKVReceiver):
return
# Register staging room bootstrap info for staging handler
+ self.chunk_staging_infos = []
if (
self.kv_mgr.enable_staging
and self.kv_mgr._staging_ctx.allocator is not None
):
- self.chunk_staging_infos = []
self.kv_mgr.register_staging_room_bootstrap(
self.bootstrap_room, self.bootstrap_infos, self
)
diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py
index 1f4c947fa..1a54a2e04 100644
--- a/python/sglang/srt/disaggregation/prefill.py
+++ b/python/sglang/srt/disaggregation/prefill.py
@@ -143,11 +143,34 @@ class PrefillBootstrapQueue:
self.scheduler.tp_worker.model_runner.effective_max_total_num_tokens
)
self.transfer_backend = transfer_backend
- if envs.SGLANG_DISAGG_STAGING_BUFFER.get() and self.is_mla_backend:
- raise RuntimeError(
- "SGLANG_DISAGG_STAGING_BUFFER is designed for non-MLA models "
- "(e.g. GQA, MHA). MLA models should not set this flag."
- )
+ if envs.SGLANG_DISAGG_STAGING_BUFFER.get():
+ if self.is_mla_backend:
+ raise RuntimeError(
+ "SGLANG_DISAGG_STAGING_BUFFER is designed for non-MLA models "
+ "(e.g. GQA, MHA). MLA models should not set this flag."
+ )
+ server_args = self.scheduler.server_args
+ page_size = self.scheduler.token_to_kv_pool_allocator.page_size
+ cps = server_args.chunked_prefill_size or 8192
+ # Staging slices each send into a fixed page-aligned grid, so an
+ # unbounded (-1) or non-page-aligned chunk size has no valid grid.
+ if cps <= 0 or cps % page_size != 0:
+ raise RuntimeError(
+ f"SGLANG_DISAGG_STAGING_BUFFER requires a positive "
+ f"chunked_prefill_size that is a multiple of page_size "
+ f"({page_size}); got {server_args.chunked_prefill_size}."
+ )
+ if self.pp_size > 1:
+ # Staging writer accounting has no pp dimension.
+ raise RuntimeError(
+ "SGLANG_DISAGG_STAGING_BUFFER does not support pp_size > 1."
+ )
+ if server_args.enable_prefill_context_parallel:
+ # CP rewrites index_slice per rank, breaking the chunk grid.
+ raise RuntimeError(
+ "SGLANG_DISAGG_STAGING_BUFFER does not support "
+ "prefill context parallelism."
+ )
self.kv_manager = self._init_kv_manager()
def _init_kv_manager(self) -> CommonKVManager:
diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py
index 20d47e37e..0d3a1b42f 100644
--- a/python/sglang/srt/environ.py
+++ b/python/sglang/srt/environ.py
@@ -477,7 +477,6 @@ class Envs:
SGLANG_HUGEPAGE_SIZE = EnvStr("")
# Staging buffer for heterogeneous TP KV transfer
SGLANG_DISAGG_STAGING_BUFFER = EnvBool(False)
- SGLANG_DISAGG_STAGING_BUFFER_SIZE_MB = EnvInt(64)
SGLANG_DISAGG_STAGING_POOL_SIZE_MB = EnvInt(4096)
# TODO(yangminl): remove SGLANG_STAGING_USE_TORCH and the torch fallback in
# staging_buffer.py once Triton kernels are fully validated in production.
diff --git a/test/registered/disaggregation/test_disaggregation_different_tp.py b/test/registered/disaggregation/test_disaggregation_different_tp.py
index c3f47ea65..a1972895d 100644
--- a/test/registered/disaggregation/test_disaggregation_different_tp.py
+++ b/test/registered/disaggregation/test_disaggregation_different_tp.py
@@ -336,9 +336,10 @@ class TestDisaggregationMooncakeMHADecodeLargerTP(PDDisaggregationServerBase):
self.assertGreater(metrics["score"], 0.60)
+# The prefill staging buffer is auto-sized to one chunk (chunked_prefill_size);
+# the decode ring is sized manually.
STAGING_ENV = {
"SGLANG_DISAGG_STAGING_BUFFER": "1",
- "SGLANG_DISAGG_STAGING_BUFFER_SIZE_MB": "64",
"SGLANG_DISAGG_STAGING_POOL_SIZE_MB": "1024",
}
diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py
index 5b205275e..47b369414 100644
--- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py
+++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py
@@ -791,6 +791,7 @@ class TestNixlStaging(CustomTestCase):
def test_do_staging_transfer_requeues_when_allocation_not_ready(self):
mgr = self._make_manager()
+ mgr._staging_ctx = PrefillStagingContext()
strategy = MagicMock()
strategy.check_ready.return_value = (False, 0, -1, 0, -1)
kv_chunk = TransferKVChunk(