fa3/fa4: sync-free for all backends and phases (#29589)

Co-authored-by: ronhuafeng <ronhuafeng@users.noreply.github.com>
This commit is contained in:
Liangsheng Yin
2026-07-13 15:09:57 -05:00
committed by GitHub
co-authored by ronhuafeng
parent 92fc692411
commit 2ab531cfcf
3 changed files with 108 additions and 35 deletions
@@ -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,