[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
|
hidden_states: torch.Tensor
|
||||||
max_top_k: int = 1
|
max_top_k: int = 1
|
||||||
uniform_top_k_value: Optional[int] = None
|
uniform_top_k_value: Optional[int] = None
|
||||||
reserved_seq_lens_cpu: Optional[torch.Tensor] = None
|
nxt_kv_lens_cpu: Optional[torch.Tensor] = None
|
||||||
reserved_seq_lens_sum: Optional[int] = None
|
nxt_kv_lens_sum: Optional[int] = None
|
||||||
_prepare_batch_seq_lens_cpu_buf: Optional[torch.Tensor] = None
|
_prepare_batch_seq_lens_cpu_buf: Optional[torch.Tensor] = None
|
||||||
_prepare_cur_kv_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
|
_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
|
page_size = batch.token_to_kv_pool_allocator.page_size
|
||||||
nxt_kv_lens_cpu_t = self._prepare_nxt_kv_lens_cpu_buf[:bs]
|
nxt_kv_lens_cpu_t = self._prepare_nxt_kv_lens_cpu_buf[:bs]
|
||||||
committed_seq_lens_sum = 0
|
committed_seq_lens_sum = 0
|
||||||
reserved_seq_lens_sum = 0
|
nxt_kv_lens_sum = 0
|
||||||
num_needed_tokens = 0
|
num_needed_tokens = 0
|
||||||
max_top_k = 1
|
max_top_k = 1
|
||||||
uniform_top_k_value = None
|
uniform_top_k_value = None
|
||||||
@@ -153,17 +153,25 @@ class DFlashDraftInputV2(SpecInput):
|
|||||||
for i, req in enumerate(batch.reqs):
|
for i, req in enumerate(batch.reqs):
|
||||||
committed_len = int(req.kv_committed_len)
|
committed_len = int(req.kv_committed_len)
|
||||||
# Read the allocation watermark from the req object like EAGLE.
|
# Read the allocation watermark from the req object like EAGLE.
|
||||||
cur_alloc_len = int(req.kv.kv_allocated_len)
|
cur = int(req.kv.kv_allocated_len)
|
||||||
reserved_len = max(cur_alloc_len, committed_len + 2 * block_size)
|
# 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)
|
top_k = int(req.sampling_params.top_k)
|
||||||
|
|
||||||
batch_seq_lens_cpu_t[i] = committed_len
|
batch_seq_lens_cpu_t[i] = committed_len
|
||||||
cur_kv_lens_cpu_t[i] = cur_alloc_len
|
cur_kv_lens_cpu_t[i] = cur
|
||||||
nxt_kv_lens_cpu_t[i] = reserved_len
|
nxt_kv_lens_cpu_t[i] = nxt
|
||||||
|
|
||||||
committed_seq_lens_sum += committed_len
|
committed_seq_lens_sum += committed_len
|
||||||
reserved_seq_lens_sum += reserved_len
|
nxt_kv_lens_sum += nxt
|
||||||
num_needed_tokens += reserved_len - cur_alloc_len
|
num_needed_tokens += nxt - cur
|
||||||
|
|
||||||
if top_k > max_top_k:
|
if top_k > max_top_k:
|
||||||
max_top_k = 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.
|
# Seed committed; overlap's resolve overwrites it with the published value.
|
||||||
batch.seq_lens_cpu = batch_seq_lens_cpu_t
|
batch.seq_lens_cpu = batch_seq_lens_cpu_t
|
||||||
batch.seq_lens_sum = committed_seq_lens_sum
|
batch.seq_lens_sum = committed_seq_lens_sum
|
||||||
self.reserved_seq_lens_cpu = nxt_kv_lens_cpu_t
|
self.nxt_kv_lens_cpu = nxt_kv_lens_cpu_t
|
||||||
self.reserved_seq_lens_sum = reserved_seq_lens_sum
|
self.nxt_kv_lens_sum = nxt_kv_lens_sum
|
||||||
|
|
||||||
def filter_batch(
|
def filter_batch(
|
||||||
self,
|
self,
|
||||||
new_indices: torch.Tensor,
|
new_indices: torch.Tensor,
|
||||||
new_indices_cpu: Optional[List[int]] = None,
|
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:
|
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:
|
else:
|
||||||
self.reserved_seq_lens_cpu = self.reserved_seq_lens_cpu[
|
self.nxt_kv_lens_cpu = self.nxt_kv_lens_cpu[new_indices.cpu()]
|
||||||
new_indices.cpu()
|
self.nxt_kv_lens_sum = int(self.nxt_kv_lens_cpu.sum().item())
|
||||||
]
|
|
||||||
self.reserved_seq_lens_sum = int(self.reserved_seq_lens_cpu.sum().item())
|
|
||||||
|
|
||||||
if self.future_indices is not None:
|
if self.future_indices is not None:
|
||||||
self.future_indices = self.future_indices[new_indices]
|
self.future_indices = self.future_indices[new_indices]
|
||||||
@@ -241,15 +247,15 @@ class DFlashDraftInputV2(SpecInput):
|
|||||||
self.hidden_states = self.hidden_states[new_indices]
|
self.hidden_states = self.hidden_states[new_indices]
|
||||||
|
|
||||||
def merge_batch(self, spec_info: "DFlashDraftInputV2"):
|
def merge_batch(self, spec_info: "DFlashDraftInputV2"):
|
||||||
if self.reserved_seq_lens_cpu is not None:
|
if self.nxt_kv_lens_cpu is not None:
|
||||||
assert spec_info.reserved_seq_lens_cpu is not None
|
assert spec_info.nxt_kv_lens_cpu is not None
|
||||||
self.reserved_seq_lens_cpu = torch.cat(
|
self.nxt_kv_lens_cpu = torch.cat(
|
||||||
[self.reserved_seq_lens_cpu, spec_info.reserved_seq_lens_cpu]
|
[self.nxt_kv_lens_cpu, spec_info.nxt_kv_lens_cpu]
|
||||||
)
|
)
|
||||||
self.reserved_seq_lens_sum = int(self.reserved_seq_lens_cpu.sum().item())
|
self.nxt_kv_lens_sum = int(self.nxt_kv_lens_cpu.sum().item())
|
||||||
elif spec_info.reserved_seq_lens_cpu is not None:
|
elif spec_info.nxt_kv_lens_cpu is not None:
|
||||||
self.reserved_seq_lens_cpu = spec_info.reserved_seq_lens_cpu
|
self.nxt_kv_lens_cpu = spec_info.nxt_kv_lens_cpu
|
||||||
self.reserved_seq_lens_sum = spec_info.reserved_seq_lens_sum
|
self.nxt_kv_lens_sum = spec_info.nxt_kv_lens_sum
|
||||||
|
|
||||||
if self.future_indices is not None:
|
if self.future_indices is not None:
|
||||||
assert spec_info.future_indices is not None
|
assert spec_info.future_indices is not None
|
||||||
|
|||||||
@@ -673,7 +673,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
batch_seq_lens_cpu: Optional[torch.Tensor],
|
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,
|
draft_prefix_lens: torch.Tensor,
|
||||||
out: torch.Tensor,
|
out: torch.Tensor,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -682,8 +682,8 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
(same contract as the non-compact path in forward_batch_generation)."""
|
(same contract as the non-compact path in forward_batch_generation)."""
|
||||||
if batch_seq_lens_cpu is not None:
|
if batch_seq_lens_cpu is not None:
|
||||||
self._compute_compact_draft_seq_lens_host(batch_seq_lens_cpu, out=out)
|
self._compute_compact_draft_seq_lens_host(batch_seq_lens_cpu, out=out)
|
||||||
elif reserved_seq_lens_cpu is not None:
|
elif nxt_kv_lens_cpu is not None:
|
||||||
self._compute_compact_draft_seq_lens_host(reserved_seq_lens_cpu, out=out)
|
self._compute_compact_draft_seq_lens_host(nxt_kv_lens_cpu, out=out)
|
||||||
else:
|
else:
|
||||||
# Last resort: the legacy blocking D2H copy.
|
# Last resort: the legacy blocking D2H copy.
|
||||||
out.copy_(draft_prefix_lens)
|
out.copy_(draft_prefix_lens)
|
||||||
@@ -1637,7 +1637,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
draft_prefix_lens = self._compute_compact_draft_seq_lens(prefix_lens)
|
draft_prefix_lens = self._compute_compact_draft_seq_lens(prefix_lens)
|
||||||
self._fill_compact_seq_lens_cpu_bound(
|
self._fill_compact_seq_lens_cpu_bound(
|
||||||
batch_seq_lens_cpu=batch.seq_lens_cpu,
|
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,
|
draft_prefix_lens=draft_prefix_lens,
|
||||||
out=seq_lens_cpu,
|
out=seq_lens_cpu,
|
||||||
)
|
)
|
||||||
@@ -1661,10 +1661,10 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
seq_lens_cpu.copy_(batch.seq_lens_cpu)
|
seq_lens_cpu.copy_(batch.seq_lens_cpu)
|
||||||
seq_lens_cpu.add_(block_size)
|
seq_lens_cpu.add_(block_size)
|
||||||
draft_seq_lens_sum = int(seq_lens_cpu.sum())
|
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.
|
# GPU-only backend: reserved is a safe over-estimate.
|
||||||
seq_lens_cpu.copy_(draft_input.reserved_seq_lens_cpu)
|
seq_lens_cpu.copy_(draft_input.nxt_kv_lens_cpu)
|
||||||
draft_seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
|
draft_seq_lens_sum = int(draft_input.nxt_kv_lens_sum)
|
||||||
else:
|
else:
|
||||||
seq_lens_cpu.copy_(prefix_lens.to("cpu", dtype=torch.int32))
|
seq_lens_cpu.copy_(prefix_lens.to("cpu", dtype=torch.int32))
|
||||||
draft_seq_lens_sum = int(prefix_lens.sum().item())
|
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
|
verify_host_seq_lens = seq_lens_cpu_backup + block_size
|
||||||
batch.seq_lens_cpu = verify_host_seq_lens
|
batch.seq_lens_cpu = verify_host_seq_lens
|
||||||
batch.seq_lens_sum = int(verify_host_seq_lens.sum())
|
batch.seq_lens_sum = int(verify_host_seq_lens.sum())
|
||||||
elif draft_input.reserved_seq_lens_cpu is not None:
|
elif draft_input.nxt_kv_lens_cpu is not None:
|
||||||
batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu
|
batch.seq_lens_cpu = draft_input.nxt_kv_lens_cpu
|
||||||
batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
|
batch.seq_lens_sum = int(draft_input.nxt_kv_lens_sum)
|
||||||
|
|
||||||
verify_forward_batch, _ = verify_input.prepare_for_verify(
|
verify_forward_batch, _ = verify_input.prepare_for_verify(
|
||||||
batch, self.target_worker
|
batch, self.target_worker
|
||||||
|
|||||||
@@ -342,9 +342,9 @@ class DraftBlockProposer:
|
|||||||
if batch.seq_lens_cpu is not None:
|
if batch.seq_lens_cpu is not None:
|
||||||
draft_seq_lens_cpu = batch.seq_lens_cpu + gamma
|
draft_seq_lens_cpu = batch.seq_lens_cpu + gamma
|
||||||
draft_seq_lens_sum = int(draft_seq_lens_cpu.sum())
|
draft_seq_lens_sum = int(draft_seq_lens_cpu.sum())
|
||||||
elif draft_input.reserved_seq_lens_cpu is not None:
|
elif draft_input.nxt_kv_lens_cpu is not None:
|
||||||
draft_seq_lens_cpu = draft_input.reserved_seq_lens_cpu
|
draft_seq_lens_cpu = draft_input.nxt_kv_lens_cpu
|
||||||
draft_seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
|
draft_seq_lens_sum = int(draft_input.nxt_kv_lens_sum)
|
||||||
else:
|
else:
|
||||||
raise RuntimeError("DSpark decode expected batch.seq_lens_cpu, got None")
|
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:
|
if seq_lens_cpu_backup is not None:
|
||||||
batch.seq_lens_cpu = seq_lens_cpu_backup + verify_w
|
batch.seq_lens_cpu = seq_lens_cpu_backup + verify_w
|
||||||
batch.seq_lens_sum = int(batch.seq_lens_cpu.sum())
|
batch.seq_lens_sum = int(batch.seq_lens_cpu.sum())
|
||||||
elif draft_input.reserved_seq_lens_cpu is not None:
|
elif draft_input.nxt_kv_lens_cpu is not None:
|
||||||
batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu
|
batch.seq_lens_cpu = draft_input.nxt_kv_lens_cpu
|
||||||
batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
|
batch.seq_lens_sum = int(draft_input.nxt_kv_lens_sum)
|
||||||
|
|
||||||
result = self._forward_prepared_verify(
|
result = self._forward_prepared_verify(
|
||||||
batch=batch,
|
batch=batch,
|
||||||
|
|||||||
@@ -256,10 +256,8 @@ class TestFilterBatchHostIndices(CustomTestCase):
|
|||||||
|
|
||||||
def make():
|
def make():
|
||||||
info = DFlashDraftInputV2.create_idle_input(device=torch.device("cpu"))
|
info = DFlashDraftInputV2.create_idle_input(device=torch.device("cpu"))
|
||||||
info.reserved_seq_lens_cpu = torch.tensor(
|
info.nxt_kv_lens_cpu = torch.tensor([10, 20, 30, 40], dtype=torch.int32)
|
||||||
[10, 20, 30, 40], dtype=torch.int32
|
info.nxt_kv_lens_sum = 100
|
||||||
)
|
|
||||||
info.reserved_seq_lens_sum = 100
|
|
||||||
info.future_indices = torch.tensor([5, 6, 7, 8])
|
info.future_indices = torch.tensor([5, 6, 7, 8])
|
||||||
return info
|
return info
|
||||||
|
|
||||||
@@ -270,8 +268,8 @@ class TestFilterBatchHostIndices(CustomTestCase):
|
|||||||
new_indices=torch.tensor(keep),
|
new_indices=torch.tensor(keep),
|
||||||
new_indices_cpu=keep,
|
new_indices_cpu=keep,
|
||||||
)
|
)
|
||||||
torch.testing.assert_close(a.reserved_seq_lens_cpu, b.reserved_seq_lens_cpu)
|
torch.testing.assert_close(a.nxt_kv_lens_cpu, b.nxt_kv_lens_cpu)
|
||||||
self.assertEqual(a.reserved_seq_lens_sum, b.reserved_seq_lens_sum)
|
self.assertEqual(a.nxt_kv_lens_sum, b.nxt_kv_lens_sum)
|
||||||
torch.testing.assert_close(a.future_indices, b.future_indices)
|
torch.testing.assert_close(a.future_indices, b.future_indices)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user