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_context_len + self.page_size - 1
|
||||
) // self.page_size
|
||||
# Opt out of the seq_lens_cpu D2H only for dflash/dspark (their workers
|
||||
# adapted to the GPU-only relay); EAGLE/MTP/standalone/non-spec keep the
|
||||
# CPU mirror.
|
||||
self.needs_cpu_seq_lens = not SpeculativeAlgorithm.from_string(
|
||||
model_runner.server_args.speculative_algorithm
|
||||
).is_dflash_family()
|
||||
# Page table is built on-device (build_trtllm_mha_page_table) and the
|
||||
# tree-mask scratch is preallocated (get_verify_buffers_to_fill_after_draft),
|
||||
# so no seq_lens_cpu / seq_lens_sum D2H sync is ever needed.
|
||||
self.needs_cpu_seq_lens = False
|
||||
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
||||
self.skip_prefill = skip_prefill
|
||||
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 = (
|
||||
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
|
||||
@@ -589,12 +590,16 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
seqlens_in_batch = forward_batch.seq_lens
|
||||
batch_size = forward_batch.batch_size
|
||||
device = seqlens_in_batch.device
|
||||
# Eager path needs a host int for dynamic page-table sizing: the CPU
|
||||
# mirror when published, else a local D2H (not the overlap hot path).
|
||||
seq_lens_cpu = (
|
||||
forward_batch.seq_lens_cpu
|
||||
if forward_batch.seq_lens_cpu is not None
|
||||
else seqlens_in_batch.cpu()
|
||||
# Eager (non-cuda-graph) path: max_seq_len_k only feeds Python-side
|
||||
# page-table slicing and the scheduler_metadata heuristic -- never the
|
||||
# kernel. Use the CPU mirror when published; otherwise the static
|
||||
# max_context_len bound (over-wide page table, slightly suboptimal
|
||||
# split heuristic, but no D2H sync).
|
||||
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():
|
||||
@@ -604,7 +609,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata.cache_seqlens_int32 = (
|
||||
seqlens_in_batch + (self.speculative_step_id + 1)
|
||||
).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
|
||||
)
|
||||
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
|
||||
# tree; build plain metadata (padded output is discarded).
|
||||
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(
|
||||
0, batch_size + 1, dtype=torch.int32, device=device
|
||||
)
|
||||
@@ -636,7 +641,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
else:
|
||||
metadata.cache_seqlens_int32 = (seqlens_in_batch).to(torch.int32)
|
||||
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(
|
||||
0,
|
||||
batch_size * self.topk + 1,
|
||||
@@ -686,7 +691,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
else:
|
||||
# Normal Decode
|
||||
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(
|
||||
0, batch_size + 1, dtype=torch.int32, device=device
|
||||
)
|
||||
@@ -753,7 +758,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
).to(torch.int32)
|
||||
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
||||
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(
|
||||
0,
|
||||
@@ -776,7 +781,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
else:
|
||||
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_k = seq_lens_cpu.max().item()
|
||||
metadata.max_seq_len_k = eager_max_k
|
||||
metadata.cu_seqlens_q = torch.arange(
|
||||
0,
|
||||
batch_size * self.speculative_num_draft_tokens + 1,
|
||||
@@ -876,7 +881,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
include_draft_extend_v2=True
|
||||
):
|
||||
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(
|
||||
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)
|
||||
)
|
||||
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
|
||||
|
||||
# Setup local attention if enabled
|
||||
@@ -924,7 +931,12 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
self._pa_swa_prefill_lens[
|
||||
forward_batch.req_pool_indices[:batch_size]
|
||||
] = 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:
|
||||
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 = {
|
||||
"cache_seqlens": torch.zeros(
|
||||
max_bs, dtype=torch.int32, device=self.device
|
||||
@@ -2143,9 +2167,11 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
dtype=torch.int32,
|
||||
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(
|
||||
max_bs * self.speculative_num_draft_tokens,
|
||||
self.max_context_len,
|
||||
self.max_context_len + self.speculative_num_draft_tokens,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
),
|
||||
@@ -2384,13 +2410,19 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
|
||||
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
|
||||
def _host_max_seq_len(
|
||||
seq_lens_cpu: Optional[torch.Tensor], seq_lens: torch.Tensor
|
||||
) -> int:
|
||||
"""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 --
|
||||
not the dflash hot path (topk=1, device-side build)."""
|
||||
D2H. Only for paths that accept a host sync (currently prefill-aware
|
||||
SWA decode replay)."""
|
||||
src = seq_lens_cpu if seq_lens_cpu is not None else seq_lens.cpu()
|
||||
return src.max().item()
|
||||
|
||||
@@ -2473,8 +2505,11 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
# metadata.cu_seqlens_q already set in capture
|
||||
# metadata.cu_seqlens_k is not needed
|
||||
|
||||
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
|
||||
)
|
||||
max_seq_pages = (
|
||||
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.cache_seqlens_int32.copy_(seq_lens)
|
||||
# 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_k[1:].copy_(
|
||||
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.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_(
|
||||
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
||||
)
|
||||
|
||||
@@ -391,7 +391,13 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
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
|
||||
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
|
||||
@@ -411,7 +417,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
req_pool_indices=buffers.req_pool_indices,
|
||||
seq_lens=buffers.seq_lens,
|
||||
seq_lens_sum=seq_lens_sum,
|
||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
encoder_lens=None,
|
||||
# per-step write target; out_cache_loc is frozen at prepare() time.
|
||||
out_cache_loc=buffers.out_cache_loc[:num_tokens],
|
||||
@@ -646,10 +652,15 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
||||
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 bs != raw_bs:
|
||||
buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||
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
|
||||
# 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.bs = bs
|
||||
self.raw_num_tokens = num_tokens
|
||||
self.seq_lens_sum = (
|
||||
forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
|
||||
)
|
||||
seq_lens_sum = forward_batch.seq_lens_sum
|
||||
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)
|
||||
|
||||
@@ -693,7 +705,9 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
||||
batch size."""
|
||||
runner = self.runners[step]
|
||||
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_num_tokens = self.raw_num_tokens
|
||||
logits_output = LogitsProcessorOutput(
|
||||
|
||||
@@ -266,6 +266,20 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
||||
tree_mask_buf, position_buf = (
|
||||
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,
|
||||
position,
|
||||
@@ -279,7 +293,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
||||
top_scores_index,
|
||||
draft_tokens,
|
||||
batch.seq_lens,
|
||||
batch.seq_lens_sum,
|
||||
seq_lens_sum,
|
||||
self.topk,
|
||||
self.speculative_num_steps,
|
||||
self.speculative_num_draft_tokens,
|
||||
|
||||
Reference in New Issue
Block a user