Support FlashAttention3 page_size > 1 and topk > 1 case with paged attn and spec decode (#7725)
This commit is contained in:
@@ -376,6 +376,7 @@ class EAGLEWorker(TpModelWorker):
|
||||
if self.page_size == 1:
|
||||
for req in batch.reqs:
|
||||
req.kv_allocated_len += self.speculative_num_steps * self.topk
|
||||
# TODO: We only need self.speculative_num_steps - 1 * topk cache loc
|
||||
out_cache_loc, token_to_kv_pool_state_backup = alloc_token_slots(
|
||||
batch.tree_cache,
|
||||
num_seqs * self.speculative_num_steps * self.topk,
|
||||
@@ -403,21 +404,13 @@ class EAGLEWorker(TpModelWorker):
|
||||
# "x" means speculative draft tokens
|
||||
# "." means padded tokens
|
||||
|
||||
# TODO(lmzheng): The current implementation is still a fake support
|
||||
# for page size > 1. In the `assign_draft_cache_locs` below,
|
||||
# we directly move the indices instead of the real kv cache.
|
||||
# This only works when the kernel backend runs with page size = 1.
|
||||
# If the kernel backend runs with page size > 1, we need to
|
||||
# duplicate the real KV cache. The overhead of duplicating KV
|
||||
# cache seems okay because the draft KV cache only has one layer.
|
||||
# see a related copy operation in MHATokenToKVPool::move_kv_cache.
|
||||
|
||||
(
|
||||
prefix_lens,
|
||||
seq_lens,
|
||||
last_loc,
|
||||
self.num_new_pages_per_topk,
|
||||
self.extend_lens,
|
||||
last_page_lens,
|
||||
) = get_last_loc_large_page_size_large_top_k(
|
||||
batch.req_to_token_pool.req_to_token,
|
||||
batch.req_pool_indices,
|
||||
@@ -427,9 +420,9 @@ class EAGLEWorker(TpModelWorker):
|
||||
self.page_size,
|
||||
)
|
||||
prefix_lens_cpu = batch.seq_lens_cpu
|
||||
last_page_lens = prefix_lens_cpu % self.page_size
|
||||
last_page_lens_cpu = prefix_lens_cpu % self.page_size
|
||||
num_new_pages_per_topk = (
|
||||
last_page_lens + self.speculative_num_steps + self.page_size - 1
|
||||
last_page_lens_cpu + self.speculative_num_steps + self.page_size - 1
|
||||
) // self.page_size
|
||||
seq_lens_cpu = (
|
||||
prefix_lens_cpu // self.page_size * self.page_size
|
||||
@@ -450,6 +443,20 @@ class EAGLEWorker(TpModelWorker):
|
||||
)
|
||||
)
|
||||
|
||||
if self.page_size > 1 and self.topk > 1:
|
||||
last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0)
|
||||
duplicate_cache_len = torch.sum(last_page_lens_cpu).item() * (self.topk - 1)
|
||||
target_cache_loc = torch.zeros(
|
||||
duplicate_cache_len, dtype=torch.int32, device=self.device
|
||||
)
|
||||
source_cache_loc = torch.zeros(
|
||||
duplicate_cache_len, dtype=torch.int32, device=self.device
|
||||
)
|
||||
else:
|
||||
# When source_cache_loc is not needed, simply skip
|
||||
duplicate_cache_len = 0
|
||||
source_cache_loc, target_cache_loc, last_page_lens_cumsum = None, None, None
|
||||
|
||||
assign_draft_cache_locs[(num_seqs,)](
|
||||
batch.req_pool_indices,
|
||||
batch.req_to_token_pool.req_to_token,
|
||||
@@ -457,16 +464,25 @@ class EAGLEWorker(TpModelWorker):
|
||||
self.extend_lens,
|
||||
self.num_new_pages_per_topk,
|
||||
out_cache_loc,
|
||||
source_cache_loc,
|
||||
target_cache_loc,
|
||||
last_page_lens_cumsum,
|
||||
duplicate_cache_len,
|
||||
batch.req_to_token_pool.req_to_token.shape[1],
|
||||
self.topk,
|
||||
self.speculative_num_steps,
|
||||
self.page_size,
|
||||
next_power_of_2(num_seqs),
|
||||
next_power_of_2(self.speculative_num_steps),
|
||||
next_power_of_2(self.speculative_num_steps + self.page_size),
|
||||
)
|
||||
|
||||
if self.page_size > 1 and self.topk > 1:
|
||||
if duplicate_cache_len > 0:
|
||||
self.draft_model_runner.token_to_kv_pool.move_kv_cache(
|
||||
target_cache_loc, source_cache_loc
|
||||
)
|
||||
# Remove padded slots
|
||||
# TODO: We only need self.speculative_num_steps - 1 cache loc
|
||||
out_cache_loc = out_cache_loc[
|
||||
: num_seqs * self.topk * self.speculative_num_steps
|
||||
]
|
||||
@@ -581,7 +597,7 @@ class EAGLEWorker(TpModelWorker):
|
||||
)
|
||||
if self.hot_token_id is not None:
|
||||
topk_index = self.hot_token_id[topk_index]
|
||||
|
||||
# TODO: We only need self.speculative_num_steps - 1 cache loc
|
||||
out_cache_loc = out_cache_loc.reshape(
|
||||
forward_batch.batch_size, self.topk, self.speculative_num_steps
|
||||
)
|
||||
@@ -1056,4 +1072,11 @@ def get_last_loc_large_page_size_large_top_k(
|
||||
prefix_lens,
|
||||
)
|
||||
|
||||
return prefix_lens, seq_lens, last_loc, num_new_pages_per_topk, extend_lens
|
||||
return (
|
||||
prefix_lens,
|
||||
seq_lens,
|
||||
last_loc,
|
||||
num_new_pages_per_topk,
|
||||
extend_lens,
|
||||
last_page_lens,
|
||||
)
|
||||
|
||||
@@ -147,6 +147,10 @@ def assign_draft_cache_locs(
|
||||
extend_lens,
|
||||
num_new_pages_per_topk,
|
||||
out_cache_loc,
|
||||
source_cache_loc,
|
||||
target_cache_loc,
|
||||
last_page_lens_cumsum,
|
||||
duplicate_cache_len: tl.constexpr,
|
||||
pool_len: tl.constexpr,
|
||||
topk: tl.constexpr,
|
||||
speculative_num_steps: tl.constexpr,
|
||||
@@ -175,44 +179,73 @@ def assign_draft_cache_locs(
|
||||
mask = copy_offset < copy_len
|
||||
data = tl.load(out_cache_ptr + copy_offset, mask=mask)
|
||||
tl.store(token_pool + kv_start + copy_offset, data, mask=mask)
|
||||
|
||||
if page_size == 1 or topk == 1:
|
||||
return
|
||||
|
||||
# Part 2: Copy the indices for the last partial page
|
||||
prefix_len = tl.load(seq_lens + pid)
|
||||
last_page_len = prefix_len % page_size
|
||||
offsets = tl.arange(0, page_size)
|
||||
mask = offsets < last_page_len
|
||||
num_new_pages_per_topk_ = tl.load(num_new_pages_per_topk + pid)
|
||||
prefix_base = token_pool + prefix_len - last_page_len
|
||||
|
||||
for topk_id in range(topk):
|
||||
value = tl.load(prefix_base + offsets, mask=mask)
|
||||
tl.store(
|
||||
prefix_base + topk_id * num_new_pages_per_topk_ * page_size + offsets,
|
||||
value,
|
||||
mask=mask,
|
||||
)
|
||||
|
||||
# Part 3: Remove the padding in out_cache_loc
|
||||
iter_offest = tl.arange(0, iter_upper)
|
||||
for topk_id in range(topk):
|
||||
indices = tl.load(
|
||||
prefix_base
|
||||
+ topk_id * num_new_pages_per_topk_ * page_size
|
||||
+ last_page_len
|
||||
+ iter_offest,
|
||||
mask=iter_offest < speculative_num_steps,
|
||||
)
|
||||
tl.store(
|
||||
out_cache_loc
|
||||
+ pid * topk * speculative_num_steps
|
||||
+ topk_id * speculative_num_steps
|
||||
+ iter_offest,
|
||||
indices,
|
||||
mask=iter_offest < speculative_num_steps,
|
||||
)
|
||||
if page_size != 1 and topk != 1 and duplicate_cache_len > 0:
|
||||
# Part 2: Copy indices into source_cache_loc and target_cache_loc
|
||||
# Expected output: src:[8,9,10,8,9,10...] tgt:[16,17,18,24,25,26...]
|
||||
prefix_len = tl.load(seq_lens + pid)
|
||||
last_page_len = prefix_len % page_size
|
||||
offsets = tl.arange(0, page_size)
|
||||
mask = offsets < last_page_len
|
||||
num_new_pages_per_topk_ = tl.load(num_new_pages_per_topk + pid)
|
||||
prefix_base = token_pool + prefix_len - last_page_len
|
||||
src_indices = tl.load(prefix_base + offsets, mask=mask)
|
||||
last_page_lens_cumsum_ = tl.load(last_page_lens_cumsum + pid)
|
||||
# Skip the first one since no copy is needed
|
||||
for topk_id in range(1, topk):
|
||||
tl.store(
|
||||
source_cache_loc
|
||||
+ (topk - 1) * (last_page_lens_cumsum_ - last_page_len)
|
||||
+ (topk_id - 1) * last_page_len
|
||||
+ offsets,
|
||||
src_indices,
|
||||
mask=mask,
|
||||
)
|
||||
tgt_indices = tl.load(
|
||||
prefix_base + topk_id * num_new_pages_per_topk_ * page_size + offsets,
|
||||
mask=mask,
|
||||
)
|
||||
tl.store(
|
||||
target_cache_loc
|
||||
+ (topk - 1) * (last_page_lens_cumsum_ - last_page_len)
|
||||
+ (topk_id - 1) * last_page_len
|
||||
+ offsets,
|
||||
tgt_indices,
|
||||
mask=mask,
|
||||
)
|
||||
# Part 3: Copy and remove the used indices for duplication
|
||||
# speculative_num_steps=5, page_size=4, num_new_pages_per_topk_=2, last_page_len=1
|
||||
# - xxxxx .. | - xxxxx .. |
|
||||
# topk=0 topk=1
|
||||
# "-" means prefix tokens
|
||||
# "x" means speculative draft tokens
|
||||
# "." means padded tokens
|
||||
# we only want to copy the "x" part.
|
||||
iter_offset = tl.arange(0, iter_upper)
|
||||
for topk_id in range(topk):
|
||||
mask_upper = iter_offset < (speculative_num_steps + last_page_len)
|
||||
mask_lower = iter_offset >= last_page_len
|
||||
combined_mask = mask_upper & mask_lower
|
||||
indices = tl.load(
|
||||
prefix_base
|
||||
+ topk_id * num_new_pages_per_topk_ * page_size
|
||||
+ iter_offset,
|
||||
mask=combined_mask,
|
||||
other=0,
|
||||
)
|
||||
# Shift from previous batches
|
||||
ptr_offset = pid * speculative_num_steps * topk
|
||||
# Subtract last_page_len to fill the gap of duplicated last page tokens.
|
||||
# For example, token pool is (1, 2, 3, 4 ,5) and last page is 1,
|
||||
# we write 2, 3, 4 to the front of out_cache_loc.
|
||||
tl.store(
|
||||
out_cache_loc
|
||||
+ ptr_offset
|
||||
+ topk_id * speculative_num_steps
|
||||
- last_page_len
|
||||
+ iter_offset,
|
||||
indices,
|
||||
mask=combined_mask,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
|
||||
Reference in New Issue
Block a user