[Spec] Capture DFLASH draft greedy sampling inside the draft decode cuda graph (#29395)
This commit is contained in:
@@ -802,12 +802,28 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
):
|
):
|
||||||
kwargs["input_embeds"] = self.buffers.input_embeds[:num_tokens]
|
kwargs["input_embeds"] = self.buffers.input_embeds[:num_tokens]
|
||||||
|
|
||||||
return forward(
|
out = forward(
|
||||||
forward_batch.input_ids,
|
forward_batch.input_ids,
|
||||||
forward_batch.positions,
|
forward_batch.positions,
|
||||||
forward_batch,
|
forward_batch,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
dflash_sampler = getattr(
|
||||||
|
self.model_runner, "dflash_draft_sampler", None
|
||||||
|
)
|
||||||
|
if dflash_sampler is not None:
|
||||||
|
# Must be captured here, or replay leaves a stale output buffer
|
||||||
|
# the worker would read as valid tokens -- fail loudly instead.
|
||||||
|
if (
|
||||||
|
not isinstance(out, LogitsProcessorOutput)
|
||||||
|
or out.hidden_states is None
|
||||||
|
):
|
||||||
|
raise RuntimeError(
|
||||||
|
"DFLASH draft sampler set but the draft forward has no "
|
||||||
|
"hidden_states to capture into the graph."
|
||||||
|
)
|
||||||
|
dflash_sampler(out.hidden_states)
|
||||||
|
return out
|
||||||
|
|
||||||
self.deepep_adapter.capture(is_extend_in_batch=False)
|
self.deepep_adapter.capture(is_extend_in_batch=False)
|
||||||
canary_ctx = (
|
canary_ctx = (
|
||||||
|
|||||||
@@ -59,6 +59,40 @@ def _get_fused_kv_materialize_helper():
|
|||||||
return _FusedKVMaterializeHelper
|
return _FusedKVMaterializeHelper
|
||||||
|
|
||||||
|
|
||||||
|
class _DflashDraftSampler:
|
||||||
|
"""Capture-safe greedy argmax over the target LM head, run inside the draft
|
||||||
|
cuda graph so the draft sampling is captured and counted in fwd_occupancy.
|
||||||
|
DFLASH's draft has no head of its own; it borrows the target `lm_head`.
|
||||||
|
tp=1 / no-added-vocab only; TP>1 stays eager in the worker.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, *, weight, block_size, num_org, org_vocab_start, max_bs):
|
||||||
|
self.weight = weight
|
||||||
|
self.block_size = int(block_size)
|
||||||
|
self.num_org = int(num_org)
|
||||||
|
self.org_vocab_start = int(org_vocab_start)
|
||||||
|
# Proposed draft tokens: written in-graph, read by the worker after replay.
|
||||||
|
self.out = torch.empty(
|
||||||
|
(int(max_bs) * (self.block_size - 1),),
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=weight.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
def __call__(self, hidden_states):
|
||||||
|
# draft tokens are block positions 1: (pos 0 is the seeded bonus token)
|
||||||
|
bs = hidden_states.shape[0] // self.block_size
|
||||||
|
hs = hidden_states.view(bs, self.block_size, -1)[:, 1:, :].reshape(
|
||||||
|
-1, hidden_states.shape[-1]
|
||||||
|
)
|
||||||
|
if hs.dtype != self.weight.dtype:
|
||||||
|
hs = hs.to(self.weight.dtype)
|
||||||
|
logits = torch.matmul(hs, self.weight[: self.num_org].T)
|
||||||
|
tokens = torch.argmax(logits, dim=-1).to(torch.long)
|
||||||
|
if self.org_vocab_start:
|
||||||
|
tokens += self.org_vocab_start
|
||||||
|
self.out[: tokens.shape[0]].copy_(tokens)
|
||||||
|
|
||||||
|
|
||||||
class DFlashWorkerV2(BaseSpecWorker):
|
class DFlashWorkerV2(BaseSpecWorker):
|
||||||
"""DFLASH speculative decoding worker (spec-v2).
|
"""DFLASH speculative decoding worker (spec-v2).
|
||||||
|
|
||||||
@@ -158,6 +192,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
)
|
)
|
||||||
set_global_server_args_for_scheduler(saved_server_args)
|
set_global_server_args_for_scheduler(saved_server_args)
|
||||||
self.draft_model_runner = self._draft_worker.model_runner
|
self.draft_model_runner = self._draft_worker.model_runner
|
||||||
|
self._draft_sampler = None
|
||||||
# Keep the same alias that other spec-v2 workers expose.
|
# Keep the same alias that other spec-v2 workers expose.
|
||||||
self._draft_worker.draft_runner = self.draft_model_runner
|
self._draft_worker.draft_runner = self.draft_model_runner
|
||||||
self.draft_model = self.draft_model_runner.model
|
self.draft_model = self.draft_model_runner.model
|
||||||
@@ -305,10 +340,50 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
"memory is available after target backend initialization.",
|
"memory is available after target backend initialization.",
|
||||||
available_mem,
|
available_mem,
|
||||||
)
|
)
|
||||||
|
if capture_decode_cuda_graph:
|
||||||
|
# Must run before capture so the draft graph folds the head in.
|
||||||
|
self._draft_sampler = self._maybe_build_draft_sampler()
|
||||||
|
self.draft_model_runner.dflash_draft_sampler = self._draft_sampler
|
||||||
self._draft_worker.init_cuda_graphs(
|
self._draft_worker.init_cuda_graphs(
|
||||||
capture_decode_cuda_graph=capture_decode_cuda_graph
|
capture_decode_cuda_graph=capture_decode_cuda_graph
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _maybe_build_draft_sampler(self):
|
||||||
|
def _eager(reason):
|
||||||
|
if self.tp_rank == 0:
|
||||||
|
logger.info("DFLASH draft greedy head kept eager (reason=%s).", reason)
|
||||||
|
return None
|
||||||
|
|
||||||
|
if get_tp_group().world_size != 1:
|
||||||
|
return _eager("tp>1")
|
||||||
|
if self.block_size <= 1:
|
||||||
|
return _eager("block_size<=1")
|
||||||
|
target_model = self._target_worker.model_runner.model
|
||||||
|
lm_head = getattr(target_model, "lm_head", None)
|
||||||
|
if lm_head is None or not hasattr(lm_head, "weight"):
|
||||||
|
return _eager("no target lm_head")
|
||||||
|
if not torch.is_floating_point(lm_head.weight):
|
||||||
|
# Quantized lm_head (FP8/INT) would break the static matmul.
|
||||||
|
return _eager("quantized lm_head")
|
||||||
|
if not hasattr(lm_head, "shard_indices"):
|
||||||
|
num_org = int(lm_head.weight.shape[0])
|
||||||
|
org_vocab_start = 0
|
||||||
|
else:
|
||||||
|
shard = lm_head.shard_indices
|
||||||
|
if int(shard.num_added_elements) != 0:
|
||||||
|
return _eager("added vocab")
|
||||||
|
num_org = int(shard.num_org_elements)
|
||||||
|
org_vocab_start = int(shard.org_vocab_start_index)
|
||||||
|
if self.tp_rank == 0:
|
||||||
|
logger.info("DFLASH draft greedy head folded into the draft cuda graph.")
|
||||||
|
return _DflashDraftSampler(
|
||||||
|
weight=lm_head.weight,
|
||||||
|
block_size=self.block_size,
|
||||||
|
num_org=num_org,
|
||||||
|
org_vocab_start=org_vocab_start,
|
||||||
|
max_bs=max(self.server_args.cuda_graph_config.decode.bs),
|
||||||
|
)
|
||||||
|
|
||||||
def _init_fused_kv_helper(self) -> None:
|
def _init_fused_kv_helper(self) -> None:
|
||||||
"""Initialize the fused KV materialization helper with pre-stacked weights."""
|
"""Initialize the fused KV materialization helper with pre-stacked weights."""
|
||||||
try:
|
try:
|
||||||
@@ -1480,18 +1555,24 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
)
|
)
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
draft_logits_output = self.draft_model_runner.forward(
|
draft_out = self.draft_model_runner.forward(forward_batch)
|
||||||
forward_batch
|
draft_logits_output = draft_out.logits_output
|
||||||
).logits_output
|
|
||||||
|
|
||||||
draft_hidden = draft_logits_output.hidden_states
|
if self._draft_sampler is not None and draft_out.can_run_graph:
|
||||||
if draft_hidden is None:
|
draft_next = self._draft_sampler.out[
|
||||||
raise RuntimeError("DFLASH draft model returned no hidden states.")
|
: bs * (int(self.block_size) - 1)
|
||||||
draft_hidden = draft_hidden.view(bs, int(self.block_size), -1)
|
].view(bs, int(self.block_size) - 1)
|
||||||
draft_next = self._greedy_sample_from_vocab_parallel_head(
|
else:
|
||||||
hidden_states=draft_hidden[:, 1:, :].reshape(-1, draft_hidden.shape[-1]),
|
draft_hidden = draft_logits_output.hidden_states
|
||||||
lm_head=lm_head,
|
if draft_hidden is None:
|
||||||
).view(bs, int(self.block_size) - 1)
|
raise RuntimeError("DFLASH draft model returned no hidden states.")
|
||||||
|
draft_hidden = draft_hidden.view(bs, int(self.block_size), -1)
|
||||||
|
draft_next = self._greedy_sample_from_vocab_parallel_head(
|
||||||
|
hidden_states=draft_hidden[:, 1:, :].reshape(
|
||||||
|
-1, draft_hidden.shape[-1]
|
||||||
|
),
|
||||||
|
lm_head=lm_head,
|
||||||
|
).view(bs, int(self.block_size) - 1)
|
||||||
|
|
||||||
draft_tokens = self._draft_block_tokens_buf[:bs]
|
draft_tokens = self._draft_block_tokens_buf[:bs]
|
||||||
draft_tokens[:, 0].copy_(block_ids[:, 0])
|
draft_tokens[:, 0].copy_(block_ids[:, 0])
|
||||||
|
|||||||
Reference in New Issue
Block a user