[Spec] Page-align the DFLASH decode KV reservation (#35265)
This commit is contained in:
@@ -43,8 +43,8 @@ class DFlashDraftInputV2(SpecInput):
|
||||
hidden_states: torch.Tensor
|
||||
max_top_k: int = 1
|
||||
uniform_top_k_value: Optional[int] = None
|
||||
reserved_seq_lens_cpu: Optional[torch.Tensor] = None
|
||||
reserved_seq_lens_sum: Optional[int] = None
|
||||
nxt_kv_lens_cpu: Optional[torch.Tensor] = None
|
||||
nxt_kv_lens_sum: Optional[int] = None
|
||||
_prepare_batch_seq_lens_cpu_buf: Optional[torch.Tensor] = None
|
||||
_prepare_cur_kv_lens_cpu_buf: Optional[torch.Tensor] = None
|
||||
_prepare_nxt_kv_lens_cpu_buf: Optional[torch.Tensor] = None
|
||||
@@ -145,7 +145,7 @@ class DFlashDraftInputV2(SpecInput):
|
||||
page_size = batch.token_to_kv_pool_allocator.page_size
|
||||
nxt_kv_lens_cpu_t = self._prepare_nxt_kv_lens_cpu_buf[:bs]
|
||||
committed_seq_lens_sum = 0
|
||||
reserved_seq_lens_sum = 0
|
||||
nxt_kv_lens_sum = 0
|
||||
num_needed_tokens = 0
|
||||
max_top_k = 1
|
||||
uniform_top_k_value = None
|
||||
@@ -153,17 +153,25 @@ class DFlashDraftInputV2(SpecInput):
|
||||
for i, req in enumerate(batch.reqs):
|
||||
committed_len = int(req.kv_committed_len)
|
||||
# Read the allocation watermark from the req object like EAGLE.
|
||||
cur_alloc_len = int(req.kv.kv_allocated_len)
|
||||
reserved_len = max(cur_alloc_len, committed_len + 2 * block_size)
|
||||
cur = int(req.kv.kv_allocated_len)
|
||||
# Whole-page accounting (same as eagle_prepare_for_decode): the
|
||||
# paged allocator hands out full pages, so an unaligned reserve
|
||||
# strands the tail of the last page -- allocated but never recorded.
|
||||
nxt = max(
|
||||
cur,
|
||||
(committed_len + 2 * block_size + page_size - 1)
|
||||
// page_size
|
||||
* page_size,
|
||||
)
|
||||
top_k = int(req.sampling_params.top_k)
|
||||
|
||||
batch_seq_lens_cpu_t[i] = committed_len
|
||||
cur_kv_lens_cpu_t[i] = cur_alloc_len
|
||||
nxt_kv_lens_cpu_t[i] = reserved_len
|
||||
cur_kv_lens_cpu_t[i] = cur
|
||||
nxt_kv_lens_cpu_t[i] = nxt
|
||||
|
||||
committed_seq_lens_sum += committed_len
|
||||
reserved_seq_lens_sum += reserved_len
|
||||
num_needed_tokens += reserved_len - cur_alloc_len
|
||||
nxt_kv_lens_sum += nxt
|
||||
num_needed_tokens += nxt - cur
|
||||
|
||||
if top_k > max_top_k:
|
||||
max_top_k = top_k
|
||||
@@ -213,22 +221,20 @@ class DFlashDraftInputV2(SpecInput):
|
||||
# Seed committed; overlap's resolve overwrites it with the published value.
|
||||
batch.seq_lens_cpu = batch_seq_lens_cpu_t
|
||||
batch.seq_lens_sum = committed_seq_lens_sum
|
||||
self.reserved_seq_lens_cpu = nxt_kv_lens_cpu_t
|
||||
self.reserved_seq_lens_sum = reserved_seq_lens_sum
|
||||
self.nxt_kv_lens_cpu = nxt_kv_lens_cpu_t
|
||||
self.nxt_kv_lens_sum = nxt_kv_lens_sum
|
||||
|
||||
def filter_batch(
|
||||
self,
|
||||
new_indices: torch.Tensor,
|
||||
new_indices_cpu: Optional[List[int]] = None,
|
||||
):
|
||||
if self.reserved_seq_lens_cpu is not None:
|
||||
if self.nxt_kv_lens_cpu is not None:
|
||||
if new_indices_cpu is not None:
|
||||
self.reserved_seq_lens_cpu = self.reserved_seq_lens_cpu[new_indices_cpu]
|
||||
self.nxt_kv_lens_cpu = self.nxt_kv_lens_cpu[new_indices_cpu]
|
||||
else:
|
||||
self.reserved_seq_lens_cpu = self.reserved_seq_lens_cpu[
|
||||
new_indices.cpu()
|
||||
]
|
||||
self.reserved_seq_lens_sum = int(self.reserved_seq_lens_cpu.sum().item())
|
||||
self.nxt_kv_lens_cpu = self.nxt_kv_lens_cpu[new_indices.cpu()]
|
||||
self.nxt_kv_lens_sum = int(self.nxt_kv_lens_cpu.sum().item())
|
||||
|
||||
if self.future_indices is not None:
|
||||
self.future_indices = self.future_indices[new_indices]
|
||||
@@ -241,15 +247,15 @@ class DFlashDraftInputV2(SpecInput):
|
||||
self.hidden_states = self.hidden_states[new_indices]
|
||||
|
||||
def merge_batch(self, spec_info: "DFlashDraftInputV2"):
|
||||
if self.reserved_seq_lens_cpu is not None:
|
||||
assert spec_info.reserved_seq_lens_cpu is not None
|
||||
self.reserved_seq_lens_cpu = torch.cat(
|
||||
[self.reserved_seq_lens_cpu, spec_info.reserved_seq_lens_cpu]
|
||||
if self.nxt_kv_lens_cpu is not None:
|
||||
assert spec_info.nxt_kv_lens_cpu is not None
|
||||
self.nxt_kv_lens_cpu = torch.cat(
|
||||
[self.nxt_kv_lens_cpu, spec_info.nxt_kv_lens_cpu]
|
||||
)
|
||||
self.reserved_seq_lens_sum = int(self.reserved_seq_lens_cpu.sum().item())
|
||||
elif spec_info.reserved_seq_lens_cpu is not None:
|
||||
self.reserved_seq_lens_cpu = spec_info.reserved_seq_lens_cpu
|
||||
self.reserved_seq_lens_sum = spec_info.reserved_seq_lens_sum
|
||||
self.nxt_kv_lens_sum = int(self.nxt_kv_lens_cpu.sum().item())
|
||||
elif spec_info.nxt_kv_lens_cpu is not None:
|
||||
self.nxt_kv_lens_cpu = spec_info.nxt_kv_lens_cpu
|
||||
self.nxt_kv_lens_sum = spec_info.nxt_kv_lens_sum
|
||||
|
||||
if self.future_indices is not None:
|
||||
assert spec_info.future_indices is not None
|
||||
|
||||
@@ -673,7 +673,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
self,
|
||||
*,
|
||||
batch_seq_lens_cpu: Optional[torch.Tensor],
|
||||
reserved_seq_lens_cpu: Optional[torch.Tensor],
|
||||
nxt_kv_lens_cpu: Optional[torch.Tensor],
|
||||
draft_prefix_lens: torch.Tensor,
|
||||
out: torch.Tensor,
|
||||
) -> None:
|
||||
@@ -682,8 +682,8 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
(same contract as the non-compact path in forward_batch_generation)."""
|
||||
if batch_seq_lens_cpu is not None:
|
||||
self._compute_compact_draft_seq_lens_host(batch_seq_lens_cpu, out=out)
|
||||
elif reserved_seq_lens_cpu is not None:
|
||||
self._compute_compact_draft_seq_lens_host(reserved_seq_lens_cpu, out=out)
|
||||
elif nxt_kv_lens_cpu is not None:
|
||||
self._compute_compact_draft_seq_lens_host(nxt_kv_lens_cpu, out=out)
|
||||
else:
|
||||
# Last resort: the legacy blocking D2H copy.
|
||||
out.copy_(draft_prefix_lens)
|
||||
@@ -1637,7 +1637,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
draft_prefix_lens = self._compute_compact_draft_seq_lens(prefix_lens)
|
||||
self._fill_compact_seq_lens_cpu_bound(
|
||||
batch_seq_lens_cpu=batch.seq_lens_cpu,
|
||||
reserved_seq_lens_cpu=draft_input.reserved_seq_lens_cpu,
|
||||
nxt_kv_lens_cpu=draft_input.nxt_kv_lens_cpu,
|
||||
draft_prefix_lens=draft_prefix_lens,
|
||||
out=seq_lens_cpu,
|
||||
)
|
||||
@@ -1661,10 +1661,10 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
seq_lens_cpu.copy_(batch.seq_lens_cpu)
|
||||
seq_lens_cpu.add_(block_size)
|
||||
draft_seq_lens_sum = int(seq_lens_cpu.sum())
|
||||
elif draft_input.reserved_seq_lens_cpu is not None:
|
||||
elif draft_input.nxt_kv_lens_cpu is not None:
|
||||
# GPU-only backend: reserved is a safe over-estimate.
|
||||
seq_lens_cpu.copy_(draft_input.reserved_seq_lens_cpu)
|
||||
draft_seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
|
||||
seq_lens_cpu.copy_(draft_input.nxt_kv_lens_cpu)
|
||||
draft_seq_lens_sum = int(draft_input.nxt_kv_lens_sum)
|
||||
else:
|
||||
seq_lens_cpu.copy_(prefix_lens.to("cpu", dtype=torch.int32))
|
||||
draft_seq_lens_sum = int(prefix_lens.sum().item())
|
||||
@@ -1740,9 +1740,9 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
verify_host_seq_lens = seq_lens_cpu_backup + block_size
|
||||
batch.seq_lens_cpu = verify_host_seq_lens
|
||||
batch.seq_lens_sum = int(verify_host_seq_lens.sum())
|
||||
elif draft_input.reserved_seq_lens_cpu is not None:
|
||||
batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu
|
||||
batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
|
||||
elif draft_input.nxt_kv_lens_cpu is not None:
|
||||
batch.seq_lens_cpu = draft_input.nxt_kv_lens_cpu
|
||||
batch.seq_lens_sum = int(draft_input.nxt_kv_lens_sum)
|
||||
|
||||
verify_forward_batch, _ = verify_input.prepare_for_verify(
|
||||
batch, self.target_worker
|
||||
|
||||
@@ -342,9 +342,9 @@ class DraftBlockProposer:
|
||||
if batch.seq_lens_cpu is not None:
|
||||
draft_seq_lens_cpu = batch.seq_lens_cpu + gamma
|
||||
draft_seq_lens_sum = int(draft_seq_lens_cpu.sum())
|
||||
elif draft_input.reserved_seq_lens_cpu is not None:
|
||||
draft_seq_lens_cpu = draft_input.reserved_seq_lens_cpu
|
||||
draft_seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
|
||||
elif draft_input.nxt_kv_lens_cpu is not None:
|
||||
draft_seq_lens_cpu = draft_input.nxt_kv_lens_cpu
|
||||
draft_seq_lens_sum = int(draft_input.nxt_kv_lens_sum)
|
||||
else:
|
||||
raise RuntimeError("DSpark decode expected batch.seq_lens_cpu, got None")
|
||||
|
||||
|
||||
@@ -256,9 +256,9 @@ class TargetVerifyExecutor:
|
||||
if seq_lens_cpu_backup is not None:
|
||||
batch.seq_lens_cpu = seq_lens_cpu_backup + verify_w
|
||||
batch.seq_lens_sum = int(batch.seq_lens_cpu.sum())
|
||||
elif draft_input.reserved_seq_lens_cpu is not None:
|
||||
batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu
|
||||
batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
|
||||
elif draft_input.nxt_kv_lens_cpu is not None:
|
||||
batch.seq_lens_cpu = draft_input.nxt_kv_lens_cpu
|
||||
batch.seq_lens_sum = int(draft_input.nxt_kv_lens_sum)
|
||||
|
||||
result = self._forward_prepared_verify(
|
||||
batch=batch,
|
||||
|
||||
Reference in New Issue
Block a user