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_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,