[misc] Add CI-only guards for the FutureMap seq_lens relay (#30471)

This commit is contained in:
Liangsheng Yin
2026-07-08 15:29:53 -07:00
committed by GitHub
parent 096551eed6
commit 8f9307736a
+33 -8
View File
@@ -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():