From 8f9307736ac3a569cb2cec7004c74de79001fa7c Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 8 Jul 2026 15:29:53 -0700 Subject: [PATCH] [misc] Add CI-only guards for the FutureMap seq_lens relay (#30471) --- python/sglang/srt/managers/overlap_utils.py | 41 +++++++++++++++++---- 1 file changed, 33 insertions(+), 8 deletions(-) diff --git a/python/sglang/srt/managers/overlap_utils.py b/python/sglang/srt/managers/overlap_utils.py index 82b6a6025..6e6e89e80 100644 --- a/python/sglang/srt/managers/overlap_utils.py +++ b/python/sglang/srt/managers/overlap_utils.py @@ -146,16 +146,21 @@ class FutureMap: self.needs_cpu_seq_lens = needs_cpu_seq_lens self.req_pool_size = req_to_token_pool.req_to_token.shape[0] - self.output_tokens_buf = ( - torch.full((self.req_pool_size,), -1, dtype=torch.int64, device=self.device) - if _DEBUG_ASSERT - else torch.empty( + if _DEBUG_ASSERT: + # Poisoned init: every row must be written before its first gather. + self.output_tokens_buf = torch.full( + (self.req_pool_size,), -1, dtype=torch.int64, device=self.device + ) + self.new_seq_lens_buf = torch.full( + (self.req_pool_size,), -1, dtype=torch.int64, device=self.device + ) + else: + self.output_tokens_buf = torch.empty( + (self.req_pool_size,), dtype=torch.int64, device=self.device + ) + self.new_seq_lens_buf = torch.empty( (self.req_pool_size,), dtype=torch.int64, device=self.device ) - ) - self.new_seq_lens_buf = torch.empty( - (self.req_pool_size,), dtype=torch.int64, device=self.device - ) # Pinned host copy of new_seq_lens_buf + private stream for fwd-prepare # D2H pulls (gated only on publish, off the schedule stream). CUDA-only: # recovers occupancy lost to the WAR barrier (also CUDA-only); other @@ -172,6 +177,9 @@ class FutureMap: self._forward_buf_initialized = False self.publish_ready = None # lazy device.Event(); only spec_v2 needs it + # Debug consume-once state: armed by a recording publish, consumed by + # resolve; arm/consume strictly alternate across all batch interleavings. + self._publish_fresh = False def _lazy_init_forward_buf(self, payload: RelayPayload): # Local import (see decide_needs_cpu_seq_lens): keep module-level deps leaf. @@ -285,6 +293,11 @@ class FutureMap: if fi is None: return if self.publish_ready is not None: + if _DEBUG_ASSERT: + # Consume-once: every event wait must be re-armed by a fresh + # forward publish; a stale consume means a publish went missing. + assert self._publish_fresh, "resolve without a fresh forward publish" + self._publish_fresh = False if _is_hip: # Temporary workaround: Event.wait() regresses TPOT on AMD MI355. self.publish_ready.synchronize() @@ -297,11 +310,18 @@ class FutureMap: # skip the .cpu() D2H. Downstream takes the GPU-only path. batch.seq_lens_cpu = None batch.seq_lens_sum = None + if _DEBUG_ASSERT: + # Poison consumed rows: each row must be re-published/seeded + # before the next resolve gathers it (safe here: the forward's + # re-publish is fenced behind this stream via wait_stream). + _assert_nonneg_and_invalidate(batch.seq_lens, self.new_seq_lens_buf, fi) return if self.fwd_prepare_d2h_stream is None or self.publish_ready is None: batch.seq_lens_cpu = batch.seq_lens.cpu() # bootstrap / non-CUDA batch.seq_lens_sum = int(batch.seq_lens_cpu.sum()) + if _DEBUG_ASSERT: + _assert_nonneg_and_invalidate(batch.seq_lens, self.new_seq_lens_buf, fi) return # Mechanism: don't sync the schedule stream; gate a private stream on the @@ -314,6 +334,10 @@ class FutureMap: # FIXME: fi == batch.req_pool_indices; unify future_indices and req_pool_indices. batch.seq_lens_cpu = self.new_seq_lens_cpu_pinned[batch.req_pool_indices_cpu] batch.seq_lens_sum = int(batch.seq_lens_cpu.sum()) + if _DEBUG_ASSERT: + # After the D2H copy completed (synchronize above), so the pinned + # mirror is not poisoned. + _assert_nonneg_and_invalidate(batch.seq_lens, self.new_seq_lens_buf, fi) def publish(self, future_indices: torch.Tensor, new_seq_lens: torch.Tensor) -> None: indices = future_indices @@ -331,6 +355,7 @@ class FutureMap: # seeding) cannot drop the in-flight forward's fence. device_module.current_stream().wait_event(self.publish_ready) self.publish_ready.record() + self._publish_fresh = True def stash(self, future_indices: torch.Tensor, payload: RelayPayload) -> None: if self.spec_algo.is_ngram():