[FlashInfer] Avoid D2H sync for sliding-window lengths (#32218)
Co-authored-by: llilian73 <204300658+llilian73@users.noreply.github.com> Co-authored-by: hnyls2002 <lsyincs@gmail.com> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
co-authored by
llilian73
hnyls2002
Liangsheng Yin
parent
0f18d389b4
commit
b24c8f10e7
@@ -1964,14 +1964,19 @@ class FlashInferIndicesUpdaterPrefill:
|
|||||||
)
|
)
|
||||||
if prefix_lens is None:
|
if prefix_lens is None:
|
||||||
num_accept_tokens = getattr(spec_info, "num_accept_tokens", None)
|
num_accept_tokens = getattr(spec_info, "num_accept_tokens", None)
|
||||||
|
# Spec verify keeps its query block outside seq_lens, so an unset
|
||||||
|
# prefix means the whole sequence is already-cached prefix.
|
||||||
|
prefix_is_full_seq = num_accept_tokens is None
|
||||||
prefix_lens = (
|
prefix_lens = (
|
||||||
seq_lens
|
seq_lens
|
||||||
if num_accept_tokens is None
|
if prefix_is_full_seq
|
||||||
else seq_lens
|
else seq_lens
|
||||||
- num_accept_tokens[: seq_lens.shape[0]].to(
|
- num_accept_tokens[: seq_lens.shape[0]].to(
|
||||||
device=seq_lens.device, dtype=seq_lens.dtype
|
device=seq_lens.device, dtype=seq_lens.dtype
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
prefix_is_full_seq = False
|
||||||
sliding_window_size = self.sliding_window_size
|
sliding_window_size = self.sliding_window_size
|
||||||
assert sliding_window_size is not None
|
assert sliding_window_size is not None
|
||||||
for wrapper_id in range(2):
|
for wrapper_id in range(2):
|
||||||
@@ -1997,6 +2002,13 @@ class FlashInferIndicesUpdaterPrefill:
|
|||||||
seq_lens,
|
seq_lens,
|
||||||
sliding_window_size + seq_lens - prefix_lens,
|
sliding_window_size + seq_lens - prefix_lens,
|
||||||
)
|
)
|
||||||
|
if prefix_is_full_seq and seq_lens_cpu is not None:
|
||||||
|
# prefix_lens is seq_lens, so the trim is min(seq_lens, window);
|
||||||
|
# summing the host mirror avoids draining the stream.
|
||||||
|
paged_kernel_lens_sum = int(
|
||||||
|
torch.clamp(seq_lens_cpu, max=sliding_window_size).sum()
|
||||||
|
)
|
||||||
|
else:
|
||||||
paged_kernel_lens_sum = paged_kernel_lens.sum().item()
|
paged_kernel_lens_sum = paged_kernel_lens.sum().item()
|
||||||
kv_start_idx = seq_lens - paged_kernel_lens
|
kv_start_idx = seq_lens - paged_kernel_lens
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -120,6 +120,22 @@ class TestFlashInferSWAAttentionBackendCorrectness(CustomTestCase):
|
|||||||
1,
|
1,
|
||||||
"dflash",
|
"dflash",
|
||||||
),
|
),
|
||||||
|
(
|
||||||
|
DenseAttentionCase(
|
||||||
|
name="runner_dflash_verify_swa_window_edges",
|
||||||
|
backend="flashinfer",
|
||||||
|
forward_mode=ForwardMode.TARGET_VERIFY,
|
||||||
|
num_heads=4,
|
||||||
|
num_kv_heads=4,
|
||||||
|
page_size=16,
|
||||||
|
# Straddle the window: one request below, one at, one above.
|
||||||
|
prefix_lens=(1, 4, 9),
|
||||||
|
extend_lens=(3, 3, 3),
|
||||||
|
sliding_window_size=4,
|
||||||
|
),
|
||||||
|
1,
|
||||||
|
"dflash",
|
||||||
|
),
|
||||||
)
|
)
|
||||||
SPEC_VERIFY_CUDA_GRAPH_CASES = (
|
SPEC_VERIFY_CUDA_GRAPH_CASES = (
|
||||||
(
|
(
|
||||||
|
|||||||
Reference in New Issue
Block a user