fa3/fa4: sync-free for all backends and phases (#29589)
Co-authored-by: ronhuafeng <ronhuafeng@users.noreply.github.com>
This commit is contained in:
co-authored by
ronhuafeng
parent
92fc692411
commit
2ab531cfcf
@@ -256,15 +256,16 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self.max_num_pages = (
|
self.max_num_pages = (
|
||||||
self.max_context_len + self.page_size - 1
|
self.max_context_len + self.page_size - 1
|
||||||
) // self.page_size
|
) // self.page_size
|
||||||
# Opt out of the seq_lens_cpu D2H only for dflash/dspark (their workers
|
# Page table is built on-device (build_trtllm_mha_page_table) and the
|
||||||
# adapted to the GPU-only relay); EAGLE/MTP/standalone/non-spec keep the
|
# tree-mask scratch is preallocated (get_verify_buffers_to_fill_after_draft),
|
||||||
# CPU mirror.
|
# so no seq_lens_cpu / seq_lens_sum D2H sync is ever needed.
|
||||||
self.needs_cpu_seq_lens = not SpeculativeAlgorithm.from_string(
|
self.needs_cpu_seq_lens = False
|
||||||
model_runner.server_args.speculative_algorithm
|
|
||||||
).is_dflash_family()
|
|
||||||
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
||||||
self.skip_prefill = skip_prefill
|
self.skip_prefill = skip_prefill
|
||||||
self.attn_cp_size = model_runner.attn_cp_size
|
self.attn_cp_size = model_runner.attn_cp_size
|
||||||
|
# Preallocated FULL_MASK tree-mask scratch; lets build_tree_kernel_efficient
|
||||||
|
# avoid the seq_lens_sum D2H sync (see get_verify_buffers_to_fill_after_draft).
|
||||||
|
self.cuda_graph_custom_mask = None
|
||||||
|
|
||||||
self.use_sliding_window_kv_pool = (
|
self.use_sliding_window_kv_pool = (
|
||||||
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
|
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
|
||||||
@@ -589,12 +590,16 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
seqlens_in_batch = forward_batch.seq_lens
|
seqlens_in_batch = forward_batch.seq_lens
|
||||||
batch_size = forward_batch.batch_size
|
batch_size = forward_batch.batch_size
|
||||||
device = seqlens_in_batch.device
|
device = seqlens_in_batch.device
|
||||||
# Eager path needs a host int for dynamic page-table sizing: the CPU
|
# Eager (non-cuda-graph) path: max_seq_len_k only feeds Python-side
|
||||||
# mirror when published, else a local D2H (not the overlap hot path).
|
# page-table slicing and the scheduler_metadata heuristic -- never the
|
||||||
seq_lens_cpu = (
|
# kernel. Use the CPU mirror when published; otherwise the static
|
||||||
forward_batch.seq_lens_cpu
|
# max_context_len bound (over-wide page table, slightly suboptimal
|
||||||
if forward_batch.seq_lens_cpu is not None
|
# split heuristic, but no D2H sync).
|
||||||
else seqlens_in_batch.cpu()
|
seq_lens_cpu = forward_batch.seq_lens_cpu
|
||||||
|
eager_max_k = (
|
||||||
|
seq_lens_cpu.max().item()
|
||||||
|
if seq_lens_cpu is not None
|
||||||
|
else self.max_context_len
|
||||||
)
|
)
|
||||||
|
|
||||||
if forward_batch.forward_mode.is_decode_or_idle():
|
if forward_batch.forward_mode.is_decode_or_idle():
|
||||||
@@ -604,7 +609,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.cache_seqlens_int32 = (
|
metadata.cache_seqlens_int32 = (
|
||||||
seqlens_in_batch + (self.speculative_step_id + 1)
|
seqlens_in_batch + (self.speculative_step_id + 1)
|
||||||
).to(torch.int32)
|
).to(torch.int32)
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item() + (
|
metadata.max_seq_len_k = eager_max_k + (
|
||||||
self.speculative_step_id + 1
|
self.speculative_step_id + 1
|
||||||
)
|
)
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
@@ -623,7 +628,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
# Draft-extend's idle batch (padded for DP MLP-sync) has no
|
# Draft-extend's idle batch (padded for DP MLP-sync) has no
|
||||||
# tree; build plain metadata (padded output is discarded).
|
# tree; build plain metadata (padded output is discarded).
|
||||||
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = eager_max_k
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0, batch_size + 1, dtype=torch.int32, device=device
|
0, batch_size + 1, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
@@ -636,7 +641,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
metadata.cache_seqlens_int32 = (seqlens_in_batch).to(torch.int32)
|
metadata.cache_seqlens_int32 = (seqlens_in_batch).to(torch.int32)
|
||||||
metadata.max_seq_len_q = self.topk
|
metadata.max_seq_len_q = self.topk
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = eager_max_k
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0,
|
0,
|
||||||
batch_size * self.topk + 1,
|
batch_size * self.topk + 1,
|
||||||
@@ -686,7 +691,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
# Normal Decode
|
# Normal Decode
|
||||||
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = eager_max_k
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0, batch_size + 1, dtype=torch.int32, device=device
|
0, batch_size + 1, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
@@ -753,7 +758,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
).to(torch.int32)
|
).to(torch.int32)
|
||||||
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
||||||
metadata.max_seq_len_k = (
|
metadata.max_seq_len_k = (
|
||||||
seq_lens_cpu.max().item() + self.speculative_num_draft_tokens
|
eager_max_k + self.speculative_num_draft_tokens
|
||||||
)
|
)
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0,
|
0,
|
||||||
@@ -776,7 +781,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
metadata.cache_seqlens_int32 = forward_batch.seq_lens.to(torch.int32)
|
metadata.cache_seqlens_int32 = forward_batch.seq_lens.to(torch.int32)
|
||||||
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = eager_max_k
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0,
|
0,
|
||||||
batch_size * self.speculative_num_draft_tokens + 1,
|
batch_size * self.speculative_num_draft_tokens + 1,
|
||||||
@@ -876,7 +881,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
include_draft_extend_v2=True
|
include_draft_extend_v2=True
|
||||||
):
|
):
|
||||||
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = eager_max_k
|
||||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||||
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
|
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
|
||||||
)
|
)
|
||||||
@@ -913,7 +918,9 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
torch.cumsum(extend_seq_lens, dim=0, dtype=torch.int32), (1, 0)
|
torch.cumsum(extend_seq_lens, dim=0, dtype=torch.int32), (1, 0)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
metadata.max_seq_len_q = metadata.max_seq_len_k
|
# max_seq_len_q reaches the kernel -- needs the real host max,
|
||||||
|
# not the static max_seq_len_k bound.
|
||||||
|
metadata.max_seq_len_q = max(forward_batch.extend_seq_lens_cpu)
|
||||||
metadata.cu_seqlens_q = metadata.cu_seqlens_k
|
metadata.cu_seqlens_q = metadata.cu_seqlens_k
|
||||||
|
|
||||||
# Setup local attention if enabled
|
# Setup local attention if enabled
|
||||||
@@ -924,7 +931,12 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self._pa_swa_prefill_lens[
|
self._pa_swa_prefill_lens[
|
||||||
forward_batch.req_pool_indices[:batch_size]
|
forward_batch.req_pool_indices[:batch_size]
|
||||||
] = forward_batch.seq_lens[:batch_size].to(torch.int32)
|
] = forward_batch.seq_lens[:batch_size].to(torch.int32)
|
||||||
max_pf = int(seq_lens_cpu[:batch_size].max().item())
|
if seq_lens_cpu is not None:
|
||||||
|
max_pf = int(seq_lens_cpu[:batch_size].max().item())
|
||||||
|
else:
|
||||||
|
# Ratchet needs a true upper bound; a local D2H beats
|
||||||
|
# poisoning it with max_context_len forever.
|
||||||
|
max_pf = int(forward_batch.seq_lens[:batch_size].max().item())
|
||||||
if max_pf > self._pa_swa_max_prefill_len:
|
if max_pf > self._pa_swa_max_prefill_len:
|
||||||
self._pa_swa_max_prefill_len = max_pf
|
self._pa_swa_max_prefill_len = max_pf
|
||||||
|
|
||||||
@@ -2040,6 +2052,18 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
),
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Worst-case FULL_MASK tree-mask scratch (bool). build_tree_kernel
|
||||||
|
# fills it in-place, so the GPU-only path needs no seq_lens_sum.
|
||||||
|
# Costs max_num_tokens * max_context_len bytes (can reach 100s of
|
||||||
|
# MB at long context) and is fully memset every verify step.
|
||||||
|
if not self.skip_prefill:
|
||||||
|
self.cuda_graph_custom_mask = torch.zeros(
|
||||||
|
max_num_tokens
|
||||||
|
* (self.max_context_len + self.speculative_num_draft_tokens),
|
||||||
|
dtype=torch.bool,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
self.draft_extend_metadata = {
|
self.draft_extend_metadata = {
|
||||||
"cache_seqlens": torch.zeros(
|
"cache_seqlens": torch.zeros(
|
||||||
max_bs, dtype=torch.int32, device=self.device
|
max_bs, dtype=torch.int32, device=self.device
|
||||||
@@ -2143,9 +2167,11 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
),
|
),
|
||||||
|
# Width covers the static max_seq_len_k bound + draft
|
||||||
|
# columns (checked by assert_buffer_fits at the merge).
|
||||||
"page_table": torch.zeros(
|
"page_table": torch.zeros(
|
||||||
max_bs * self.speculative_num_draft_tokens,
|
max_bs * self.speculative_num_draft_tokens,
|
||||||
self.max_context_len,
|
self.max_context_len + self.speculative_num_draft_tokens,
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
),
|
),
|
||||||
@@ -2384,13 +2410,19 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
|
|
||||||
return metadata, metadata_expand
|
return metadata, metadata_expand
|
||||||
|
|
||||||
|
def get_verify_buffers_to_fill_after_draft(self):
|
||||||
|
# Return the preallocated FULL_MASK tree-mask scratch so that
|
||||||
|
# build_tree_kernel_efficient fills it in-place and the worker never
|
||||||
|
# needs seq_lens_sum to size a dynamic allocation (no D2H sync).
|
||||||
|
return [self.cuda_graph_custom_mask, None]
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _host_max_seq_len(
|
def _host_max_seq_len(
|
||||||
seq_lens_cpu: Optional[torch.Tensor], seq_lens: torch.Tensor
|
seq_lens_cpu: Optional[torch.Tensor], seq_lens: torch.Tensor
|
||||||
) -> int:
|
) -> int:
|
||||||
"""Host-side max KV length: the CPU mirror when published, else a local
|
"""Host-side max KV length: the CPU mirror when published, else a local
|
||||||
D2H. For cold paths (topk>1, draft-extend, eager) that need a host max --
|
D2H. Only for paths that accept a host sync (currently prefill-aware
|
||||||
not the dflash hot path (topk=1, device-side build)."""
|
SWA decode replay)."""
|
||||||
src = seq_lens_cpu if seq_lens_cpu is not None else seq_lens.cpu()
|
src = seq_lens_cpu if seq_lens_cpu is not None else seq_lens.cpu()
|
||||||
return src.max().item()
|
return src.max().item()
|
||||||
|
|
||||||
@@ -2473,8 +2505,11 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
# metadata.cu_seqlens_q already set in capture
|
# metadata.cu_seqlens_q already set in capture
|
||||||
# metadata.cu_seqlens_k is not needed
|
# metadata.cu_seqlens_k is not needed
|
||||||
|
|
||||||
metadata.max_seq_len_k = self._host_max_seq_len(
|
# Tight page-table bound when the CPU mirror is free.
|
||||||
seq_lens_cpu, seq_lens
|
metadata.max_seq_len_k = (
|
||||||
|
seq_lens_cpu.max().item()
|
||||||
|
if seq_lens_cpu is not None
|
||||||
|
else self.max_context_len
|
||||||
)
|
)
|
||||||
max_seq_pages = (
|
max_seq_pages = (
|
||||||
metadata.max_seq_len_k + self.page_size - 1
|
metadata.max_seq_len_k + self.page_size - 1
|
||||||
@@ -2636,7 +2671,12 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata = self.target_verify_metadata_topk_normal[bs]
|
metadata = self.target_verify_metadata_topk_normal[bs]
|
||||||
metadata.cache_seqlens_int32.copy_(seq_lens)
|
metadata.cache_seqlens_int32.copy_(seq_lens)
|
||||||
# metadata.max_seq_len_q = self.speculative_num_draft_tokens, already set in capture
|
# metadata.max_seq_len_q = self.speculative_num_draft_tokens, already set in capture
|
||||||
metadata.max_seq_len_k = self._host_max_seq_len(seq_lens_cpu, seq_lens)
|
# Tight page-table bound when the CPU mirror is free.
|
||||||
|
metadata.max_seq_len_k = (
|
||||||
|
seq_lens_cpu.max().item()
|
||||||
|
if seq_lens_cpu is not None
|
||||||
|
else self.max_context_len
|
||||||
|
)
|
||||||
# metadata.cu_seqlens_q already set in capture
|
# metadata.cu_seqlens_q already set in capture
|
||||||
metadata.cu_seqlens_k[1:].copy_(
|
metadata.cu_seqlens_k[1:].copy_(
|
||||||
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
||||||
@@ -2722,7 +2762,12 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata = self.draft_extend_metadata[bs]
|
metadata = self.draft_extend_metadata[bs]
|
||||||
metadata.cache_seqlens_int32.copy_(seq_lens)
|
metadata.cache_seqlens_int32.copy_(seq_lens)
|
||||||
|
|
||||||
metadata.max_seq_len_k = self._host_max_seq_len(seq_lens_cpu, seq_lens)
|
# Tight page-table bound when the CPU mirror is free.
|
||||||
|
metadata.max_seq_len_k = (
|
||||||
|
seq_lens_cpu.max().item()
|
||||||
|
if seq_lens_cpu is not None
|
||||||
|
else self.max_context_len
|
||||||
|
)
|
||||||
metadata.cu_seqlens_k[1:].copy_(
|
metadata.cu_seqlens_k[1:].copy_(
|
||||||
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -391,7 +391,13 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
post_warmup_hook=post_warmup_hook,
|
post_warmup_hook=post_warmup_hook,
|
||||||
)
|
)
|
||||||
|
|
||||||
def replay(self, bs: int, seq_lens_sum: int, spec_info: EagleDraftExtendInput):
|
def replay(
|
||||||
|
self,
|
||||||
|
bs: int,
|
||||||
|
seq_lens_sum: Optional[int],
|
||||||
|
spec_info: EagleDraftExtendInput,
|
||||||
|
seq_lens_cpu: Optional[torch.Tensor],
|
||||||
|
):
|
||||||
"""Init this step's attention metadata for the prepared bucket and
|
"""Init this step's attention metadata for the prepared bucket and
|
||||||
replay its graph. Buffers must already be populated by the composite
|
replay its graph. Buffers must already be populated by the composite
|
||||||
runner's ``prepare`` (step 0) or by the previous step's in-graph chain
|
runner's ``prepare`` (step 0) or by the previous step's in-graph chain
|
||||||
@@ -411,7 +417,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
req_pool_indices=buffers.req_pool_indices,
|
req_pool_indices=buffers.req_pool_indices,
|
||||||
seq_lens=buffers.seq_lens,
|
seq_lens=buffers.seq_lens,
|
||||||
seq_lens_sum=seq_lens_sum,
|
seq_lens_sum=seq_lens_sum,
|
||||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
encoder_lens=None,
|
encoder_lens=None,
|
||||||
# per-step write target; out_cache_loc is frozen at prepare() time.
|
# per-step write target; out_cache_loc is frozen at prepare() time.
|
||||||
out_cache_loc=buffers.out_cache_loc[:num_tokens],
|
out_cache_loc=buffers.out_cache_loc[:num_tokens],
|
||||||
@@ -646,10 +652,15 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
|||||||
forward_batch.spec_info.num_accept_tokens
|
forward_batch.spec_info.num_accept_tokens
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Refresh the host mirror only when published; hand replay None
|
||||||
|
# otherwise so no consumer reads a stale buffer.
|
||||||
if forward_batch.seq_lens_cpu is not None:
|
if forward_batch.seq_lens_cpu is not None:
|
||||||
if bs != raw_bs:
|
if bs != raw_bs:
|
||||||
buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||||
buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||||
|
self.seq_lens_cpu = buffers.seq_lens_cpu
|
||||||
|
else:
|
||||||
|
self.seq_lens_cpu = None
|
||||||
|
|
||||||
# select_index[i] = i * window + num_correct_drafts[i]: the flat index
|
# select_index[i] = i * window + num_correct_drafts[i]: the flat index
|
||||||
# of request i's last accepted token. Used by the in-graph top-k gather
|
# of request i's last accepted token. Used by the in-graph top-k gather
|
||||||
@@ -681,9 +692,10 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
|||||||
self.raw_bs = raw_bs
|
self.raw_bs = raw_bs
|
||||||
self.bs = bs
|
self.bs = bs
|
||||||
self.raw_num_tokens = num_tokens
|
self.raw_num_tokens = num_tokens
|
||||||
self.seq_lens_sum = (
|
seq_lens_sum = forward_batch.seq_lens_sum
|
||||||
forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
|
if seq_lens_sum is not None:
|
||||||
)
|
seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
|
||||||
|
self.seq_lens_sum = seq_lens_sum
|
||||||
|
|
||||||
self._prepare_extra(forward_batch)
|
self._prepare_extra(forward_batch)
|
||||||
|
|
||||||
@@ -693,7 +705,9 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
|||||||
batch size."""
|
batch size."""
|
||||||
runner = self.runners[step]
|
runner = self.runners[step]
|
||||||
runner.raw_bs = self.raw_bs
|
runner.raw_bs = self.raw_bs
|
||||||
out = runner.replay(self.bs, self.seq_lens_sum, self._replay_spec_info)
|
out = runner.replay(
|
||||||
|
self.bs, self.seq_lens_sum, self._replay_spec_info, self.seq_lens_cpu
|
||||||
|
)
|
||||||
raw_bs = self.raw_bs
|
raw_bs = self.raw_bs
|
||||||
raw_num_tokens = self.raw_num_tokens
|
raw_num_tokens = self.raw_num_tokens
|
||||||
logits_output = LogitsProcessorOutput(
|
logits_output = LogitsProcessorOutput(
|
||||||
|
|||||||
@@ -266,6 +266,20 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
tree_mask_buf, position_buf = (
|
tree_mask_buf, position_buf = (
|
||||||
self.target_worker.model_runner.attn_backend.get_verify_buffers_to_fill_after_draft()
|
self.target_worker.model_runner.attn_backend.get_verify_buffers_to_fill_after_draft()
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# build_tree_kernel uses seq_lens_sum only to size the (non-preallocated)
|
||||||
|
# tree mask; over-size is safe. Skip per-iter .sum().item() D2H via UB.
|
||||||
|
seq_lens_sum = batch.seq_lens_sum
|
||||||
|
if seq_lens_sum is None:
|
||||||
|
if tree_mask_buf is None:
|
||||||
|
max_context_len = (
|
||||||
|
self.target_worker.model_runner.attn_backend.max_context_len
|
||||||
|
)
|
||||||
|
seq_lens_sum = batch.seq_lens.shape[0] * max_context_len
|
||||||
|
else:
|
||||||
|
# tree_mask_buf preallocated -> kernel ignores seq_lens_sum.
|
||||||
|
seq_lens_sum = 0
|
||||||
|
|
||||||
(
|
(
|
||||||
tree_mask,
|
tree_mask,
|
||||||
position,
|
position,
|
||||||
@@ -279,7 +293,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
top_scores_index,
|
top_scores_index,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
batch.seq_lens,
|
batch.seq_lens,
|
||||||
batch.seq_lens_sum,
|
seq_lens_sum,
|
||||||
self.topk,
|
self.topk,
|
||||||
self.speculative_num_steps,
|
self.speculative_num_steps,
|
||||||
self.speculative_num_draft_tokens,
|
self.speculative_num_draft_tokens,
|
||||||
|
|||||||
Reference in New Issue
Block a user