[Spec] Replace shared-infra dflash special-cases with capabilities (WAR barrier + seq_lens_cpu) (#29232)

This commit is contained in:
Liangsheng Yin
2026-06-28 00:02:32 -07:00
committed by GitHub
parent 97ddc66209
commit 593f5ba68f
4 changed files with 67 additions and 162 deletions
+2 -18
View File
@@ -224,10 +224,6 @@ class FutureMap:
if draft_input is None:
# FIXME(lsyin): only prefill; not compatible with mixed mode
return
if self.spec_algo.is_dflash() and getattr(
draft_input, "direct_carry_valid", False
):
return
indices = draft_input.future_indices
if indices.shape[0] == 0:
return
@@ -266,18 +262,11 @@ class FutureMap:
def resolve_seq_lens_cpu(self, batch: ScheduleBatch) -> None:
# Lazy pull from new_seq_lens_buf for spec_v2 (accept_lens not known to
# schedule). DFLASH intentionally keeps host-side lengths lagging and
# uses its carried KV allocation watermark for planning, so only the GPU
# seq_lens is resolved there. Other spec-v2 algorithms still need the CPU
# mirror for host planning; use a private D2H stream for those copies.
# schedule). The CPU mirror is gated by needs_cpu_seq_lens; backends that
# opt out take the GPU-only path below. A private D2H stream overlaps the copy.
draft_input = batch.spec_info
if draft_input is None:
return
if self.spec_algo.is_dflash() and getattr(
draft_input, "direct_carry_valid", False
):
batch.seq_lens = draft_input.new_seq_lens
return
fi = draft_input.future_indices
if fi is None:
@@ -290,11 +279,6 @@ class FutureMap:
self.publish_ready.wait()
batch.seq_lens = self.new_seq_lens_buf[fi]
if self.spec_algo.is_dflash():
# DFLASH keeps seq_lens_cpu as the lagging committed host view;
# planning/reserved host lengths live on DFlashDraftInputV2.
return
if not self.needs_cpu_seq_lens:
# GPU gather above is kept (SB.seq_lens must advance each verify);
# skip the .cpu() D2H. Downstream takes the GPU-only path.
+4 -3
View File
@@ -1487,9 +1487,10 @@ class Scheduler(
self.schedule_stream = self.device_module.Stream(priority=0)
if self.device == "cpu":
self.schedule_stream.synchronize = lambda: None # No-op for CPU
# DFLASH fences its shared req_to_token writes with verify_done /
# plan-stream deps, so the global WAR barrier only serializes plan
# overlap. TODO: generalize this global-barrier enablement policy.
# The global WAR barrier fences the scheduler's next shared-buffer write
# on the forward's read-done event. DFLASH opts out: it fences its own
# req_to_token writes with verify_done / plan-stream deps, so the global
# barrier would only serialize plan overlap without adding correctness.
self._war_barrier_enabled = (
is_cuda() or envs.SGLANG_ENABLE_WAR_BARRIER.get()
) and not self.spec_algorithm.is_dflash()
@@ -40,9 +40,7 @@ def _get_overlap_plan_stream(
class DFlashDraftInputV2(SpecInput):
"""Draft-side state carried across overlap iterations (spec-v2)."""
# Legacy Eagle-shaped fields kept only for dataclass compatibility. DFLASH
# overlap carries new_seq_lens / bonus_tokens directly in the common
# no-shape-change path; FutureMap remains the fallback for filter/merge.
# Legacy Eagle-shaped fields; DFLASH relays via FutureMap so these are unused.
topk_p: torch.Tensor
topk_index: torch.Tensor
bonus_tokens: torch.Tensor
@@ -51,14 +49,8 @@ class DFlashDraftInputV2(SpecInput):
verify_done: Optional[torch.cuda.Event] = None
max_top_k: int = 1
uniform_top_k_value: Optional[int] = None
cur_allocated_seq_lens_cpu: Optional[torch.Tensor] = None
planning_seq_lens_cpu: Optional[torch.Tensor] = None
planning_seq_lens_sum: Optional[int] = None
reserved_seq_lens_cpu: Optional[torch.Tensor] = None
reserved_seq_lens_sum: Optional[int] = None
direct_carry_valid: bool = True
_prepare_committed_kv_lens_cpu_buf: Optional[torch.Tensor] = None
_prepare_planning_kv_lens_cpu_buf: Optional[torch.Tensor] = None
_prepare_batch_seq_lens_cpu_buf: Optional[torch.Tensor] = None
_prepare_cur_kv_lens_cpu_buf: Optional[torch.Tensor] = None
_prepare_nxt_kv_lens_cpu_buf: Optional[torch.Tensor] = None
@@ -81,7 +73,7 @@ class DFlashDraftInputV2(SpecInput):
pin_memory = is_pin_memory_available(device)
def needs_cpu_alloc(buf: Optional[torch.Tensor]) -> bool:
return buf is None or buf.numel() < bs or buf.is_pinned() != pin_memory
return buf is None or buf.numel() < bs
def needs_gpu_alloc(buf: Optional[torch.Tensor]) -> bool:
return buf is None or buf.numel() < bs or str(buf.device) != str(device)
@@ -90,14 +82,10 @@ class DFlashDraftInputV2(SpecInput):
current = 0 if buf is None else int(buf.numel())
return max(bs, 32, current * 2 if current > 0 else 0)
if needs_cpu_alloc(self._prepare_committed_kv_lens_cpu_buf):
capacity = grown_capacity(self._prepare_committed_kv_lens_cpu_buf)
self._prepare_committed_kv_lens_cpu_buf = torch.empty(
(capacity,), dtype=torch.int32, device="cpu", pin_memory=pin_memory
)
self._prepare_planning_kv_lens_cpu_buf = torch.empty(
(capacity,), dtype=torch.int32, device="cpu", pin_memory=pin_memory
)
# The three CPU scratch buffers grow together; capacity is the only
# invariant (batch is int64 non-pinned, cur/nxt are int32 pinned).
if needs_cpu_alloc(self._prepare_batch_seq_lens_cpu_buf):
capacity = grown_capacity(self._prepare_batch_seq_lens_cpu_buf)
self._prepare_batch_seq_lens_cpu_buf = torch.empty(
(capacity,), dtype=torch.int64, device="cpu"
)
@@ -146,18 +134,13 @@ class DFlashDraftInputV2(SpecInput):
if bs == 0:
return
self._ensure_prepare_length_buffers(bs, batch.device)
assert self._prepare_committed_kv_lens_cpu_buf is not None
assert self._prepare_planning_kv_lens_cpu_buf is not None
assert self._prepare_batch_seq_lens_cpu_buf is not None
assert self._prepare_cur_kv_lens_cpu_buf is not None
assert self._prepare_nxt_kv_lens_cpu_buf is not None
assert self._prepare_cur_kv_lens_gpu_buf is not None
assert self._prepare_nxt_kv_lens_gpu_buf is not None
committed_kv_lens_cpu_t = self._prepare_committed_kv_lens_cpu_buf[:bs]
planning_kv_lens_cpu_t = self._prepare_planning_kv_lens_cpu_buf[:bs]
batch_seq_lens_cpu_t = self._prepare_batch_seq_lens_cpu_buf[:bs]
cur_kv_lens_cpu_t = self._prepare_cur_kv_lens_cpu_buf[:bs]
cur_allocated_seq_lens_cpu = self.cur_allocated_seq_lens_cpu
# For DFLASH, each decode step needs a fixed-size verify block.
block_size = int(get_global_server_args().speculative_num_draft_tokens)
@@ -168,7 +151,6 @@ class DFlashDraftInputV2(SpecInput):
page_size = batch.token_to_kv_pool_allocator.page_size
nxt_kv_lens_cpu_t = self._prepare_nxt_kv_lens_cpu_buf[:bs]
committed_seq_lens_sum = 0
planning_seq_lens_sum = 0
reserved_seq_lens_sum = 0
num_needed_tokens = 0
max_top_k = 1
@@ -176,24 +158,16 @@ class DFlashDraftInputV2(SpecInput):
uniform_top_k = True
for i, req in enumerate(batch.reqs):
committed_len = int(req.kv_committed_len)
if cur_allocated_seq_lens_cpu is not None and i < len(
cur_allocated_seq_lens_cpu
):
cur_alloc_len = int(cur_allocated_seq_lens_cpu[i])
else:
cur_alloc_len = int(req.kv_allocated_len)
planning_len = committed_len + block_size
# Read the allocation watermark from the req object like EAGLE.
cur_alloc_len = int(req.kv_allocated_len)
reserved_len = max(cur_alloc_len, committed_len + 2 * block_size)
top_k = int(req.sampling_params.top_k)
committed_kv_lens_cpu_t[i] = committed_len
batch_seq_lens_cpu_t[i] = committed_len
cur_kv_lens_cpu_t[i] = cur_alloc_len
planning_kv_lens_cpu_t[i] = planning_len
nxt_kv_lens_cpu_t[i] = reserved_len
committed_seq_lens_sum += committed_len
planning_seq_lens_sum += planning_len
reserved_seq_lens_sum += reserved_len
num_needed_tokens += reserved_len - cur_alloc_len
@@ -268,32 +242,19 @@ class DFlashDraftInputV2(SpecInput):
for i, req in enumerate(batch.reqs):
req.kv_allocated_len = max(req.kv_allocated_len, int(nxt_kv_lens_cpu_t[i]))
# Preserve the lagging committed CPU view on the batch and carry the
# tighter host-side planning bound separately from the full reserved
# allocator upper bound. Overlap scheduling only drifts by at most one
# DFlash block on the committed prefix lengths.
# Seed committed; overlap's resolve overwrites it with the published value.
batch.seq_lens_cpu = batch_seq_lens_cpu_t
batch.seq_lens_sum = committed_seq_lens_sum
self.planning_seq_lens_cpu = planning_kv_lens_cpu_t
self.planning_seq_lens_sum = planning_seq_lens_sum
self.reserved_seq_lens_cpu = nxt_kv_lens_cpu_t
self.reserved_seq_lens_sum = reserved_seq_lens_sum
def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
if self.cur_allocated_seq_lens_cpu is not None:
self.cur_allocated_seq_lens_cpu = self.cur_allocated_seq_lens_cpu[
new_indices.cpu()
]
if self.planning_seq_lens_cpu is not None:
self.planning_seq_lens_cpu = self.planning_seq_lens_cpu[new_indices.cpu()]
self.planning_seq_lens_sum = int(self.planning_seq_lens_cpu.sum().item())
if self.reserved_seq_lens_cpu is not None:
self.reserved_seq_lens_cpu = self.reserved_seq_lens_cpu[new_indices.cpu()]
self.reserved_seq_lens_sum = int(self.reserved_seq_lens_cpu.sum().item())
if self.future_indices is not None:
self.future_indices = self.future_indices[new_indices]
self.direct_carry_valid = False
return
self.topk_p = self.topk_p[new_indices]
@@ -303,24 +264,6 @@ class DFlashDraftInputV2(SpecInput):
self.hidden_states = self.hidden_states[new_indices]
def merge_batch(self, spec_info: "DFlashDraftInputV2"):
if self.cur_allocated_seq_lens_cpu is not None:
assert spec_info.cur_allocated_seq_lens_cpu is not None
self.cur_allocated_seq_lens_cpu = torch.cat(
[self.cur_allocated_seq_lens_cpu, spec_info.cur_allocated_seq_lens_cpu]
)
elif spec_info.cur_allocated_seq_lens_cpu is not None:
self.cur_allocated_seq_lens_cpu = spec_info.cur_allocated_seq_lens_cpu
if self.planning_seq_lens_cpu is not None:
assert spec_info.planning_seq_lens_cpu is not None
self.planning_seq_lens_cpu = torch.cat(
[self.planning_seq_lens_cpu, spec_info.planning_seq_lens_cpu]
)
self.planning_seq_lens_sum = int(self.planning_seq_lens_cpu.sum().item())
elif spec_info.planning_seq_lens_cpu is not None:
self.planning_seq_lens_cpu = spec_info.planning_seq_lens_cpu
self.planning_seq_lens_sum = spec_info.planning_seq_lens_sum
if self.reserved_seq_lens_cpu is not None:
assert spec_info.reserved_seq_lens_cpu is not None
self.reserved_seq_lens_cpu = torch.cat(
@@ -336,7 +279,6 @@ class DFlashDraftInputV2(SpecInput):
self.future_indices = torch.cat(
[self.future_indices, spec_info.future_indices]
)
self.direct_carry_valid = False
return
self.topk_p = torch.cat([self.topk_p, spec_info.topk_p], dim=0)
@@ -1212,10 +1212,8 @@ class DFlashWorkerV2(BaseSpecWorker):
self._new_seq_lens_bufs[slot][:bs],
)
def _validate_phase1_sampling_support(
self, model_worker_batch: ScheduleBatch
) -> None:
sampling_info = model_worker_batch.sampling_info
def _validate_phase1_sampling_support(self, batch: ScheduleBatch) -> None:
sampling_info = batch.sampling_info
if sampling_info is None or sampling_info.is_all_greedy:
return
@@ -1236,7 +1234,6 @@ class DFlashWorkerV2(BaseSpecWorker):
bonus_tokens: torch.Tensor,
seq_lens: torch.Tensor,
verify_done: Optional[torch.cuda.Event] = None,
cur_allocated_seq_lens_cpu: Optional[torch.Tensor] = None,
) -> DFlashDraftInputV2:
bs = int(seq_lens.numel())
device = bonus_tokens.device
@@ -1247,7 +1244,6 @@ class DFlashWorkerV2(BaseSpecWorker):
new_seq_lens=seq_lens.to(dtype=torch.int64),
hidden_states=torch.empty((bs, 0), device=device, dtype=torch.float16),
verify_done=verify_done,
cur_allocated_seq_lens_cpu=cur_allocated_seq_lens_cpu,
)
def _make_next_draft_input_decode(
@@ -1256,7 +1252,6 @@ class DFlashWorkerV2(BaseSpecWorker):
bonus_tokens: torch.Tensor,
new_seq_lens: torch.Tensor,
verify_done: Optional[torch.cuda.Event] = None,
cur_allocated_seq_lens_cpu: Optional[torch.Tensor] = None,
) -> DFlashDraftInputV2:
bs = int(new_seq_lens.numel())
device = bonus_tokens.device
@@ -1267,35 +1262,29 @@ class DFlashWorkerV2(BaseSpecWorker):
new_seq_lens=new_seq_lens.to(dtype=torch.int64),
hidden_states=torch.empty((bs, 0), device=device, dtype=torch.float16),
verify_done=verify_done,
cur_allocated_seq_lens_cpu=cur_allocated_seq_lens_cpu,
)
def forward_batch_generation(
self,
model_worker_batch: ScheduleBatch,
batch: ScheduleBatch,
on_publish=None,
) -> GenerationBatchResult:
if getattr(model_worker_batch, "return_logprob", False):
if getattr(batch, "return_logprob", False):
raise ValueError(
"DFLASH speculative decoding does not support return_logprob yet."
)
self._validate_phase1_sampling_support(model_worker_batch)
self._validate_phase1_sampling_support(batch)
if (
model_worker_batch.forward_mode.is_extend()
or model_worker_batch.is_extend_in_batch
):
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
# Target prefill: capture DFlash aux hidden states for prompt tokens.
model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL
batch_output = self.target_worker.forward_batch_generation(
model_worker_batch
)
batch.capture_hidden_mode = CaptureHiddenMode.FULL
batch_output = self.target_worker.forward_batch_generation(batch)
logits_output, next_token_ids = (
batch_output.logits_output,
batch_output.next_token_ids,
)
batch_output.new_seq_lens = model_worker_batch.seq_lens
batch_output.new_seq_lens = batch.seq_lens
if on_publish is not None:
on_publish(batch_output.new_seq_lens)
@@ -1305,10 +1294,7 @@ class DFlashWorkerV2(BaseSpecWorker):
"Make sure the target model has DFlash layers-to-capture configured."
)
if (
model_worker_batch.extend_lens is None
or model_worker_batch.prefix_lens is None
):
if batch.extend_lens is None or batch.prefix_lens is None:
raise RuntimeError(
"DFLASH expected extend_lens / prefix_lens to be populated in extend mode, "
"but got None."
@@ -1317,14 +1303,12 @@ class DFlashWorkerV2(BaseSpecWorker):
# Materialize prompt tokens into the draft KV cache immediately. This is required
# for radix cache safety (the scheduler may update radix after prefill returns).
device = next_token_ids.device
ctx_lens = torch.tensor(
model_worker_batch.extend_lens, dtype=torch.int32, device=device
)
ctx_lens = torch.tensor(batch.extend_lens, dtype=torch.int32, device=device)
draft_seq_lens = torch.tensor(
model_worker_batch.prefix_lens, dtype=torch.int32, device=device
batch.prefix_lens, dtype=torch.int32, device=device
)
if model_worker_batch.out_cache_loc is None:
if batch.out_cache_loc is None:
raise RuntimeError(
"DFLASH prefill expected out_cache_loc, but got None."
)
@@ -1332,11 +1316,11 @@ class DFlashWorkerV2(BaseSpecWorker):
self.model_runner.server_args.attention_backend,
draft_seq_lens,
ctx_lens,
int(sum(model_worker_batch.extend_lens)),
int(sum(batch.extend_lens)),
)
self._append_target_hidden_to_draft_kv_by_loc(
target_hidden=logits_output.hidden_states,
cache_loc=model_worker_batch.out_cache_loc,
cache_loc=batch.out_cache_loc,
positions=positions,
)
@@ -1345,8 +1329,7 @@ class DFlashWorkerV2(BaseSpecWorker):
batch_output.next_draft_input = self._make_next_draft_input_prefill(
bonus_tokens=next_token_ids,
seq_lens=model_worker_batch.seq_lens,
cur_allocated_seq_lens_cpu=model_worker_batch.seq_lens_cpu,
seq_lens=batch.seq_lens,
)
verify_done = torch.get_device_module(device).Event()
verify_done.record()
@@ -1354,18 +1337,16 @@ class DFlashWorkerV2(BaseSpecWorker):
return batch_output
# Decode / target-verify stage.
if model_worker_batch.spec_info is None:
model_worker_batch.spec_info = DFlashDraftInputV2.create_idle_input(
device=self.device
)
if batch.spec_info is None:
batch.spec_info = DFlashDraftInputV2.create_idle_input(device=self.device)
draft_input = model_worker_batch.spec_info
draft_input = batch.spec_info
if not isinstance(draft_input, DFlashDraftInputV2):
raise RuntimeError(
"DFLASH spec-v2 expected DFlashDraftInputV2 state on the running batch."
)
if model_worker_batch.forward_mode.is_idle():
if batch.forward_mode.is_idle():
empty_ids = torch.empty((0,), dtype=torch.int64, device=self.device)
empty_lens = torch.empty((0,), dtype=torch.int32, device=self.device)
next_draft_input = self._make_next_draft_input_decode(
@@ -1389,11 +1370,11 @@ class DFlashWorkerV2(BaseSpecWorker):
# `seq_lens` is carried over from the previous overlap iteration and may have been
# produced on another stream.
model_worker_batch.seq_lens.record_stream(
batch.seq_lens.record_stream(
torch.get_device_module(self.device).current_stream()
)
bs = len(model_worker_batch.seq_lens)
bs = len(batch.seq_lens)
device = self.device
# --- 1) Draft a fixed block with the draft model.
@@ -1415,7 +1396,7 @@ class DFlashWorkerV2(BaseSpecWorker):
assert self._draft_seq_lens_cpu_buf is not None
block_ids = self._draft_block_ids_buf[:bs]
prefix_lens = model_worker_batch.seq_lens
prefix_lens = batch.seq_lens
positions_2d = self._draft_block_positions_buf[:bs]
verify_out_cache_loc_2d = self._draft_verify_out_cache_loc_buf[:bs]
if self._use_triton_prepare_block:
@@ -1423,7 +1404,7 @@ class DFlashWorkerV2(BaseSpecWorker):
_prepare_dflash_draft_block_unchecked(
bonus_tokens=draft_input.bonus_tokens.view(-1),
prefix_lens=prefix_lens.view(-1),
req_pool_indices=model_worker_batch.req_pool_indices.view(-1),
req_pool_indices=batch.req_pool_indices.view(-1),
req_to_token=self.model_runner.req_to_token_pool.req_to_token,
block_ids_out=block_ids,
positions_out=positions_2d,
@@ -1445,7 +1426,7 @@ class DFlashWorkerV2(BaseSpecWorker):
)
end_offset = prefix_lens + block_size
verify_out_cache_loc = assign_extend_cache_locs_func(
req_pool_indices=model_worker_batch.req_pool_indices,
req_pool_indices=batch.req_pool_indices,
req_to_token=self.model_runner.req_to_token_pool.req_to_token,
start_offset=prefix_lens,
end_offset=end_offset,
@@ -1464,7 +1445,7 @@ class DFlashWorkerV2(BaseSpecWorker):
)
end_offset = prefix_lens + block_size
verify_out_cache_loc = assign_extend_cache_locs_func(
req_pool_indices=model_worker_batch.req_pool_indices,
req_pool_indices=batch.req_pool_indices,
req_to_token=self.model_runner.req_to_token_pool.req_to_token,
start_offset=prefix_lens,
end_offset=end_offset,
@@ -1491,12 +1472,12 @@ class DFlashWorkerV2(BaseSpecWorker):
)
suffix_cache_loc = self._gather_req_to_token_segments(
req_to_token=self.model_runner.req_to_token_pool.req_to_token,
req_pool_indices=model_worker_batch.req_pool_indices,
req_pool_indices=batch.req_pool_indices,
start=suffix_start,
lengths=draft_prefix_lens,
)
assign_req_to_token_pool_func(
model_worker_batch.req_pool_indices,
batch.req_pool_indices,
self.draft_model_runner.req_to_token_pool.req_to_token,
torch.zeros_like(draft_prefix_lens),
draft_prefix_lens,
@@ -1507,7 +1488,7 @@ class DFlashWorkerV2(BaseSpecWorker):
block_end = self._draft_block_end_buf[:bs]
torch.add(draft_prefix_lens, block_size, out=block_end)
assign_req_to_token_pool_func(
model_worker_batch.req_pool_indices,
batch.req_pool_indices,
self.draft_model_runner.req_to_token_pool.req_to_token,
draft_prefix_lens,
block_end,
@@ -1521,19 +1502,15 @@ class DFlashWorkerV2(BaseSpecWorker):
# Backend planning only needs a safe upper bound for the committed
# prefix lengths, not the full allocator reservation length.
draft_seq_lens = prefix_lens
if draft_input.planning_seq_lens_cpu is not None:
seq_lens_cpu.copy_(draft_input.planning_seq_lens_cpu)
draft_seq_lens_sum = int(draft_input.planning_seq_lens_sum)
if batch.seq_lens_cpu is not None:
# Host bound = committed prefix + one verify block.
seq_lens_cpu.copy_(batch.seq_lens_cpu)
seq_lens_cpu.add_(block_size)
draft_seq_lens_sum = int(seq_lens_cpu.sum())
elif draft_input.reserved_seq_lens_cpu is not None:
# GPU-only backend: reserved is a safe over-estimate.
seq_lens_cpu.copy_(draft_input.reserved_seq_lens_cpu)
draft_seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
elif model_worker_batch.seq_lens_cpu is not None:
seq_lens_cpu.copy_(model_worker_batch.seq_lens_cpu)
draft_seq_lens_sum = (
int(model_worker_batch.seq_lens_sum)
if model_worker_batch.seq_lens_sum is not None
else int(model_worker_batch.seq_lens_cpu.sum())
)
else:
seq_lens_cpu.copy_(prefix_lens.to("cpu", dtype=torch.int32))
draft_seq_lens_sum = int(prefix_lens.sum().item())
@@ -1542,7 +1519,7 @@ class DFlashWorkerV2(BaseSpecWorker):
forward_mode=ForwardMode.TARGET_VERIFY,
batch_size=bs,
input_ids=block_ids.flatten(),
req_pool_indices=model_worker_batch.req_pool_indices,
req_pool_indices=batch.req_pool_indices,
seq_lens=draft_seq_lens,
out_cache_loc=verify_out_cache_loc,
seq_lens_sum=draft_seq_lens_sum,
@@ -1591,30 +1568,32 @@ class DFlashWorkerV2(BaseSpecWorker):
capture_hidden_mode=CaptureHiddenMode.FULL,
)
model_worker_batch.out_cache_loc = verify_out_cache_loc
sampling_info = model_worker_batch.sampling_info
batch.out_cache_loc = verify_out_cache_loc
sampling_info = batch.sampling_info
need_mamba_verify_commit = hasattr(
self.target_worker.model_runner.attn_backend,
"update_mamba_state_after_mtp_verify",
)
seq_lens_pre_verify = (
model_worker_batch.seq_lens.clone() if need_mamba_verify_commit else None
batch.seq_lens.clone() if need_mamba_verify_commit else None
)
seq_lens_cpu_backup = model_worker_batch.seq_lens_cpu
seq_lens_sum_backup = model_worker_batch.seq_lens_sum
if draft_input.planning_seq_lens_cpu is not None:
model_worker_batch.seq_lens_cpu = draft_input.planning_seq_lens_cpu
model_worker_batch.seq_lens_sum = int(draft_input.planning_seq_lens_sum)
seq_lens_cpu_backup = batch.seq_lens_cpu
seq_lens_sum_backup = batch.seq_lens_sum
if seq_lens_cpu_backup is not None:
# Verify host bound = committed prefix + one verify block (matches draft).
verify_host_seq_lens = seq_lens_cpu_backup + block_size
batch.seq_lens_cpu = verify_host_seq_lens
batch.seq_lens_sum = int(verify_host_seq_lens.sum())
elif draft_input.reserved_seq_lens_cpu is not None:
model_worker_batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu
model_worker_batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu
batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
verify_forward_batch, _ = verify_input.prepare_for_verify(
model_worker_batch, self.target_worker
batch, self.target_worker
)
model_worker_batch.seq_lens_cpu = seq_lens_cpu_backup
model_worker_batch.seq_lens_sum = seq_lens_sum_backup
batch.seq_lens_cpu = seq_lens_cpu_backup
batch.seq_lens_sum = seq_lens_sum_backup
target_out = self.target_worker.forward_batch_generation(
batch=None,
@@ -1720,7 +1699,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if need_mamba_verify_commit:
assert seq_lens_pre_verify is not None
self._update_target_mamba_state_after_verify(
batch=model_worker_batch,
batch=batch,
seq_lens_pre_verify=seq_lens_pre_verify,
commit_lens=commit_lens,
)
@@ -1752,7 +1731,6 @@ class DFlashWorkerV2(BaseSpecWorker):
next_draft_input = self._make_next_draft_input_decode(
bonus_tokens=bonus,
new_seq_lens=new_seq_lens,
cur_allocated_seq_lens_cpu=draft_input.reserved_seq_lens_cpu,
)
verify_done = torch.get_device_module(device).Event()
verify_done.record()