[Spec] Disambiguate verified_id into bonus_token(s) / accept_tokens (#24724)
This commit is contained in:
@@ -12,14 +12,14 @@ Use the verb form `accept` everywhere. Don't use the past-participle form `accep
|
|||||||
| `accepted_indices` | `accept_indices` |
|
| `accepted_indices` | `accept_indices` |
|
||||||
| `accepted_token_ids` | `accept_tokens` (also see Rule 3) |
|
| `accepted_token_ids` | `accept_tokens` (also see Rule 3) |
|
||||||
|
|
||||||
## Rule 2 — The extra/bonus token is `bonus_token`
|
## Rule 2 — The extra/bonus token is `bonus_token` / `bonus_tokens`
|
||||||
|
|
||||||
The "+1" token that the target model always emits in addition to verifying drafts is the **bonus token**. Always call it `bonus_token`.
|
The "+1" token that the target model always emits in addition to verifying drafts is the **bonus token**. Use `bonus_token` / `bonus_tokens` per Rule 7.
|
||||||
|
|
||||||
| Don't | Do |
|
| Don't | Do |
|
||||||
|---|---|
|
|---|---|
|
||||||
| `verified_id` | `bonus_token` |
|
| `verified_id` / `verified_ids` | `bonus_token` / `bonus_tokens` |
|
||||||
| `output_id` (when referring to the bonus) | `bonus_token` |
|
| `output_id` / `output_ids` (when referring to the bonus) | `bonus_token` / `bonus_tokens` |
|
||||||
|
|
||||||
`req.output_ids` (the full output history of a request) is unrelated and stays as is.
|
`req.output_ids` (the full output history of a request) is unrelated and stays as is.
|
||||||
|
|
||||||
@@ -95,13 +95,15 @@ The semantic differs by scope:
|
|||||||
|
|
||||||
## Rule 7 — Singular vs plural
|
## Rule 7 — Singular vs plural
|
||||||
|
|
||||||
Plural for any non-scalar tensor (`[bs]`-shaped, flat, or multi-dim); singular only for scalars (kernel `tl.load` results, single-int locals). Applies to all spec-decoding tensors (tokens, indices, etc.). The fixed name `bonus_token` (Rule 2) is the one singular exception even though its shape is `[bs]`.
|
Plural for any non-scalar tensor (`[bs]`-shaped, flat, or multi-dim); singular only for scalars (kernel `tl.load` results, single-int locals). Applies to all spec-decoding tensors (tokens, indices, etc.).
|
||||||
|
|
||||||
```python
|
```python
|
||||||
accept_tokens: torch.Tensor # [total_accepted] flat - plural
|
accept_tokens: torch.Tensor # [total_accepted] flat - plural
|
||||||
accept_indices: torch.Tensor # [bs, num_draft_tokens] - plural
|
accept_indices: torch.Tensor # [bs, num_draft_tokens] - plural
|
||||||
draft_tokens: torch.Tensor # [bs * num_draft_tokens] flat - plural
|
draft_tokens: torch.Tensor # [bs * num_draft_tokens] flat - plural
|
||||||
|
bonus_tokens: torch.Tensor # [bs] - plural
|
||||||
accept_token = tl.load(...) # int32 scalar in a kernel iteration - singular
|
accept_token = tl.load(...) # int32 scalar in a kernel iteration - singular
|
||||||
|
bonus_token = tl.load(...) # int32 scalar inside a kernel - singular
|
||||||
```
|
```
|
||||||
|
|
||||||
## Out of scope (these names stay as is)
|
## Out of scope (these names stay as is)
|
||||||
|
|||||||
@@ -170,7 +170,7 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
|||||||
topk_p=topk_p,
|
topk_p=topk_p,
|
||||||
topk_index=topk_index,
|
topk_index=topk_index,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
verified_id=self.output_ids,
|
bonus_tokens=self.output_ids,
|
||||||
new_seq_lens=self.seq_lens,
|
new_seq_lens=self.seq_lens,
|
||||||
)
|
)
|
||||||
spec_info.prepare_for_extend(self)
|
spec_info.prepare_for_extend(self)
|
||||||
|
|||||||
@@ -362,7 +362,7 @@ def add_output_logprobs_for_spec_v1(
|
|||||||
logprobs = torch.nn.functional.log_softmax(
|
logprobs = torch.nn.functional.log_softmax(
|
||||||
logits_output.next_token_logits / temperatures, dim=-1
|
logits_output.next_token_logits / temperatures, dim=-1
|
||||||
)
|
)
|
||||||
batch_next_token_ids = res.verified_id
|
batch_next_token_ids = res.accept_tokens
|
||||||
num_tokens_per_req = [accept + 1 for accept in num_accepted_drafts_per_req_cpu]
|
num_tokens_per_req = [accept + 1 for accept in num_accepted_drafts_per_req_cpu]
|
||||||
|
|
||||||
# We should repeat top_logprobs_nums to match num_tokens_per_req.
|
# We should repeat top_logprobs_nums to match num_tokens_per_req.
|
||||||
@@ -407,7 +407,7 @@ def add_output_logprobs_for_spec_v1(
|
|||||||
# Add output logprobs to the request
|
# Add output logprobs to the request
|
||||||
pt = 0
|
pt = 0
|
||||||
next_token_logprobs = logits_output.next_token_logprobs.tolist()
|
next_token_logprobs = logits_output.next_token_logprobs.tolist()
|
||||||
verified_ids = batch_next_token_ids.tolist()
|
accept_tokens_list = batch_next_token_ids.tolist()
|
||||||
token_top_logprobs_val = logits_output.next_token_top_logprobs_val
|
token_top_logprobs_val = logits_output.next_token_top_logprobs_val
|
||||||
token_top_logprobs_idx = logits_output.next_token_top_logprobs_idx
|
token_top_logprobs_idx = logits_output.next_token_top_logprobs_idx
|
||||||
token_ids_logprobs_val = logits_output.next_token_token_ids_logprobs_val
|
token_ids_logprobs_val = logits_output.next_token_token_ids_logprobs_val
|
||||||
@@ -416,7 +416,7 @@ def add_output_logprobs_for_spec_v1(
|
|||||||
for _ in range(num_tokens):
|
for _ in range(num_tokens):
|
||||||
if req.return_logprob:
|
if req.return_logprob:
|
||||||
req.output_token_logprobs_val.append(next_token_logprobs[pt])
|
req.output_token_logprobs_val.append(next_token_logprobs[pt])
|
||||||
req.output_token_logprobs_idx.append(verified_ids[pt])
|
req.output_token_logprobs_idx.append(accept_tokens_list[pt])
|
||||||
if req.top_logprobs_num > 0:
|
if req.top_logprobs_num > 0:
|
||||||
assert (
|
assert (
|
||||||
should_top_logprobs
|
should_top_logprobs
|
||||||
|
|||||||
@@ -86,7 +86,7 @@ class FutureMap:
|
|||||||
# Get a reference for each tensor
|
# Get a reference for each tensor
|
||||||
topk_p0 = draft_input.topk_p[0]
|
topk_p0 = draft_input.topk_p[0]
|
||||||
topk_index0 = draft_input.topk_index[0]
|
topk_index0 = draft_input.topk_index[0]
|
||||||
verified_id0 = draft_input.verified_id[0]
|
bonus_token0 = draft_input.bonus_tokens[0]
|
||||||
new_seq_lens0 = draft_input.new_seq_lens[0]
|
new_seq_lens0 = draft_input.new_seq_lens[0]
|
||||||
|
|
||||||
self.topk_p_buf = torch.empty(
|
self.topk_p_buf = torch.empty(
|
||||||
@@ -99,9 +99,9 @@ class FutureMap:
|
|||||||
dtype=topk_index0.dtype,
|
dtype=topk_index0.dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
self.verified_id_buf = torch.empty(
|
self.bonus_tokens_buf = torch.empty(
|
||||||
(self.future_buffer_len, *verified_id0.shape),
|
(self.future_buffer_len, *bonus_token0.shape),
|
||||||
dtype=verified_id0.dtype,
|
dtype=bonus_token0.dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
self.new_seq_lens_buf = torch.empty(
|
self.new_seq_lens_buf = torch.empty(
|
||||||
@@ -146,7 +146,7 @@ class FutureMap:
|
|||||||
indices.record_stream(torch.get_device_module(self.device).current_stream())
|
indices.record_stream(torch.get_device_module(self.device).current_stream())
|
||||||
draft_input.topk_p = self.topk_p_buf[indices]
|
draft_input.topk_p = self.topk_p_buf[indices]
|
||||||
draft_input.topk_index = self.topk_index_buf[indices]
|
draft_input.topk_index = self.topk_index_buf[indices]
|
||||||
draft_input.verified_id = self.verified_id_buf[indices]
|
draft_input.bonus_tokens = self.bonus_tokens_buf[indices]
|
||||||
draft_input.new_seq_lens = self.new_seq_lens_buf[indices]
|
draft_input.new_seq_lens = self.new_seq_lens_buf[indices]
|
||||||
if spec_need_hidden_states():
|
if spec_need_hidden_states():
|
||||||
draft_input.hidden_states = self.hidden_states_buf[indices]
|
draft_input.hidden_states = self.hidden_states_buf[indices]
|
||||||
@@ -181,7 +181,7 @@ class FutureMap:
|
|||||||
|
|
||||||
self.topk_p_buf[intv] = draft_input.topk_p
|
self.topk_p_buf[intv] = draft_input.topk_p
|
||||||
self.topk_index_buf[intv] = draft_input.topk_index
|
self.topk_index_buf[intv] = draft_input.topk_index
|
||||||
self.verified_id_buf[intv] = draft_input.verified_id
|
self.bonus_tokens_buf[intv] = draft_input.bonus_tokens
|
||||||
self.new_seq_lens_buf[intv] = draft_input.new_seq_lens
|
self.new_seq_lens_buf[intv] = draft_input.new_seq_lens
|
||||||
if spec_need_hidden_states():
|
if spec_need_hidden_states():
|
||||||
self.hidden_states_buf[intv] = draft_input.hidden_states
|
self.hidden_states_buf[intv] = draft_input.hidden_states
|
||||||
|
|||||||
@@ -67,7 +67,7 @@ class DFlashDraftInput(SpecInput):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
# Current token to start the next DFlash block (one per request).
|
# Current token to start the next DFlash block (one per request).
|
||||||
verified_id: torch.Tensor
|
bonus_tokens: torch.Tensor
|
||||||
|
|
||||||
# Flattened context features for tokens that need to be appended into the draft cache.
|
# Flattened context features for tokens that need to be appended into the draft cache.
|
||||||
# Shape: [sum(ctx_lens), K * hidden_size], where K is the number of target-layer
|
# Shape: [sum(ctx_lens), K * hidden_size], where K is the number of target-layer
|
||||||
@@ -92,7 +92,7 @@ class DFlashDraftInput(SpecInput):
|
|||||||
old_ctx_lens = self.ctx_lens
|
old_ctx_lens = self.ctx_lens
|
||||||
old_target_hidden = self.target_hidden
|
old_target_hidden = self.target_hidden
|
||||||
|
|
||||||
self.verified_id = self.verified_id[new_indices]
|
self.bonus_tokens = self.bonus_tokens[new_indices]
|
||||||
self.ctx_lens = old_ctx_lens[new_indices]
|
self.ctx_lens = old_ctx_lens[new_indices]
|
||||||
self.draft_seq_lens = self.draft_seq_lens[new_indices]
|
self.draft_seq_lens = self.draft_seq_lens[new_indices]
|
||||||
|
|
||||||
@@ -129,7 +129,9 @@ class DFlashDraftInput(SpecInput):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def merge_batch(self, spec_info: "DFlashDraftInput"):
|
def merge_batch(self, spec_info: "DFlashDraftInput"):
|
||||||
self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], dim=0)
|
self.bonus_tokens = torch.cat(
|
||||||
|
[self.bonus_tokens, spec_info.bonus_tokens], dim=0
|
||||||
|
)
|
||||||
self.ctx_lens = torch.cat([self.ctx_lens, spec_info.ctx_lens], dim=0)
|
self.ctx_lens = torch.cat([self.ctx_lens, spec_info.ctx_lens], dim=0)
|
||||||
self.draft_seq_lens = torch.cat(
|
self.draft_seq_lens = torch.cat(
|
||||||
[self.draft_seq_lens, spec_info.draft_seq_lens], dim=0
|
[self.draft_seq_lens, spec_info.draft_seq_lens], dim=0
|
||||||
@@ -319,7 +321,7 @@ class DFlashVerifyInput(SpecInput):
|
|||||||
"""DFlash verification for greedy and non-greedy sampling.
|
"""DFlash verification for greedy and non-greedy sampling.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
new_verified_id: int64 tensor [bs] (the new current token per request)
|
new_bonus_tokens: int64 tensor [bs] (the new current token per request)
|
||||||
commit_lens: int32 tensor [bs] (how many verify-input tokens are committed)
|
commit_lens: int32 tensor [bs] (how many verify-input tokens are committed)
|
||||||
next_target_hidden: tensor [sum(commit_lens), feature_dim]
|
next_target_hidden: tensor [sum(commit_lens), feature_dim]
|
||||||
num_accepted_drafts_per_req_cpu: list[int] (accepted draft tokens per request)
|
num_accepted_drafts_per_req_cpu: list[int] (accepted draft tokens per request)
|
||||||
@@ -389,7 +391,7 @@ class DFlashVerifyInput(SpecInput):
|
|||||||
max_acc = self.draft_token_num - 1
|
max_acc = self.draft_token_num - 1
|
||||||
num_accepted_drafts_per_req_cpu: List[int] = []
|
num_accepted_drafts_per_req_cpu: List[int] = []
|
||||||
commit_lens_cpu: List[int] = []
|
commit_lens_cpu: List[int] = []
|
||||||
new_verified_list: List[int] = []
|
new_bonus_tokens_list: List[int] = []
|
||||||
|
|
||||||
for i, req in enumerate(batch.reqs):
|
for i, req in enumerate(batch.reqs):
|
||||||
acc_len = int(packed[i, max_acc].item())
|
acc_len = int(packed[i, max_acc].item())
|
||||||
@@ -409,24 +411,24 @@ class DFlashVerifyInput(SpecInput):
|
|||||||
req.grammar.accept_token(token_id)
|
req.grammar.accept_token(token_id)
|
||||||
|
|
||||||
if req.output_ids:
|
if req.output_ids:
|
||||||
new_verified_token = int(req.output_ids[-1])
|
new_bonus_token = int(req.output_ids[-1])
|
||||||
elif req.origin_input_ids:
|
elif req.origin_input_ids:
|
||||||
# If no token was appended in this verify step, keep the current token unchanged.
|
# If no token was appended in this verify step, keep the current token unchanged.
|
||||||
new_verified_token = int(req.origin_input_ids[-1])
|
new_bonus_token = int(req.origin_input_ids[-1])
|
||||||
else:
|
else:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"DFLASH verify cannot determine current token: both output_ids and origin_input_ids are empty."
|
"DFLASH verify cannot determine current token: both output_ids and origin_input_ids are empty."
|
||||||
)
|
)
|
||||||
|
|
||||||
commit_lens_cpu.append(appended)
|
commit_lens_cpu.append(appended)
|
||||||
new_verified_list.append(new_verified_token)
|
new_bonus_tokens_list.append(new_bonus_token)
|
||||||
num_accepted_drafts_per_req_cpu.append(max(0, appended - 1))
|
num_accepted_drafts_per_req_cpu.append(max(0, appended - 1))
|
||||||
req.spec_verify_ct += 1
|
req.spec_verify_ct += 1
|
||||||
req.spec_accepted_drafts += num_accepted_drafts_per_req_cpu[-1]
|
req.spec_accepted_drafts += num_accepted_drafts_per_req_cpu[-1]
|
||||||
|
|
||||||
commit_lens = torch.tensor(commit_lens_cpu, dtype=torch.int32, device=device)
|
commit_lens = torch.tensor(commit_lens_cpu, dtype=torch.int32, device=device)
|
||||||
new_verified_id = torch.tensor(
|
new_bonus_tokens = torch.tensor(
|
||||||
new_verified_list, dtype=torch.int64, device=device
|
new_bonus_tokens_list, dtype=torch.int64, device=device
|
||||||
)
|
)
|
||||||
|
|
||||||
# Free uncommitted KV cache slots and compact out_cache_loc.
|
# Free uncommitted KV cache slots and compact out_cache_loc.
|
||||||
@@ -494,7 +496,7 @@ class DFlashVerifyInput(SpecInput):
|
|||||||
logits_output.hidden_states = None
|
logits_output.hidden_states = None
|
||||||
|
|
||||||
return (
|
return (
|
||||||
new_verified_id,
|
new_bonus_tokens,
|
||||||
commit_lens,
|
commit_lens,
|
||||||
next_target_hidden,
|
next_target_hidden,
|
||||||
num_accepted_drafts_per_req_cpu,
|
num_accepted_drafts_per_req_cpu,
|
||||||
|
|||||||
@@ -570,7 +570,7 @@ class DFlashWorker:
|
|||||||
|
|
||||||
block_ids = self._draft_block_ids_buf[:bs]
|
block_ids = self._draft_block_ids_buf[:bs]
|
||||||
block_ids.fill_(int(self._mask_token_id))
|
block_ids.fill_(int(self._mask_token_id))
|
||||||
block_ids[:, 0].copy_(draft_input.verified_id.to(torch.long))
|
block_ids[:, 0].copy_(draft_input.bonus_tokens.to(torch.long))
|
||||||
|
|
||||||
noise_embedding = embed_module(block_ids)
|
noise_embedding = embed_module(block_ids)
|
||||||
input_embeds = noise_embedding.view(-1, noise_embedding.shape[-1])
|
input_embeds = noise_embedding.view(-1, noise_embedding.shape[-1])
|
||||||
@@ -1163,7 +1163,7 @@ class DFlashWorker:
|
|||||||
model_worker_batch.extend_seq_lens
|
model_worker_batch.extend_seq_lens
|
||||||
)
|
)
|
||||||
draft_input = DFlashDraftInput(
|
draft_input = DFlashDraftInput(
|
||||||
verified_id=next_token_ids.to(torch.int64),
|
bonus_tokens=next_token_ids.to(torch.int64),
|
||||||
target_hidden=logits_output.hidden_states,
|
target_hidden=logits_output.hidden_states,
|
||||||
ctx_lens=extend_seq_lens,
|
ctx_lens=extend_seq_lens,
|
||||||
draft_seq_lens=(
|
draft_seq_lens=(
|
||||||
@@ -1213,7 +1213,7 @@ class DFlashWorker:
|
|||||||
)
|
)
|
||||||
|
|
||||||
(
|
(
|
||||||
new_verified_id,
|
new_bonus_tokens,
|
||||||
commit_lens,
|
commit_lens,
|
||||||
next_target_hidden,
|
next_target_hidden,
|
||||||
num_accepted_drafts_per_req_cpu,
|
num_accepted_drafts_per_req_cpu,
|
||||||
@@ -1232,7 +1232,7 @@ class DFlashWorker:
|
|||||||
|
|
||||||
# Update draft state for the next iteration. Also materialize the committed verify tokens
|
# Update draft state for the next iteration. Also materialize the committed verify tokens
|
||||||
# into the draft KV cache immediately so radix cache entries are safe to reuse.
|
# into the draft KV cache immediately so radix cache entries are safe to reuse.
|
||||||
draft_input.verified_id = new_verified_id
|
draft_input.bonus_tokens = new_bonus_tokens
|
||||||
draft_input.target_hidden = next_target_hidden
|
draft_input.target_hidden = next_target_hidden
|
||||||
draft_input.ctx_lens = commit_lens
|
draft_input.ctx_lens = commit_lens
|
||||||
self._append_target_hidden_to_draft_kv(batch, draft_input)
|
self._append_target_hidden_to_draft_kv(batch, draft_input)
|
||||||
@@ -1249,7 +1249,7 @@ class DFlashWorker:
|
|||||||
|
|
||||||
return GenerationBatchResult(
|
return GenerationBatchResult(
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
next_token_ids=new_verified_id,
|
next_token_ids=new_bonus_tokens,
|
||||||
num_accepted_drafts=num_accepted_drafts,
|
num_accepted_drafts=num_accepted_drafts,
|
||||||
num_accepted_drafts_per_req_cpu=num_accepted_drafts_per_req_cpu,
|
num_accepted_drafts_per_req_cpu=num_accepted_drafts_per_req_cpu,
|
||||||
can_run_cuda_graph=can_run_cuda_graph,
|
can_run_cuda_graph=can_run_cuda_graph,
|
||||||
|
|||||||
@@ -250,7 +250,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
),
|
),
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
verified_id=torch.empty(0, dtype=torch.long, device=batch.device),
|
accept_tokens=torch.empty(0, dtype=torch.long, device=batch.device),
|
||||||
num_accepted_drafts_per_req_cpu=[],
|
num_accepted_drafts_per_req_cpu=[],
|
||||||
accepted_indices=torch.full(
|
accepted_indices=torch.full(
|
||||||
(0, self.spec_steps + 1),
|
(0, self.spec_steps + 1),
|
||||||
@@ -465,7 +465,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
# Free the KV cache for unaccepted tokens
|
# Free the KV cache for unaccepted tokens
|
||||||
# TODO: fuse them
|
# TODO: fuse them
|
||||||
accept_index = accept_index[accept_index != -1]
|
accept_index = accept_index[accept_index != -1]
|
||||||
verified_id = predict[accept_index]
|
accept_tokens = predict[accept_index]
|
||||||
evict_mask = torch.full_like(self.draft_token, True, dtype=torch.bool)
|
evict_mask = torch.full_like(self.draft_token, True, dtype=torch.bool)
|
||||||
evict_mask[accept_index] = False
|
evict_mask[accept_index] = False
|
||||||
num_accepted_drafts_cpu = num_accepted_drafts.cpu()
|
num_accepted_drafts_cpu = num_accepted_drafts.cpu()
|
||||||
@@ -553,7 +553,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
|
|
||||||
draft_input = EagleDraftInput(
|
draft_input = EagleDraftInput(
|
||||||
hidden_states=batch.spec_info.hidden_states[accept_index],
|
hidden_states=batch.spec_info.hidden_states[accept_index],
|
||||||
verified_id=verified_id,
|
accept_tokens=accept_tokens,
|
||||||
num_accepted_drafts=num_accepted_drafts,
|
num_accepted_drafts=num_accepted_drafts,
|
||||||
num_accepted_tokens=num_accepted_drafts + 1,
|
num_accepted_tokens=num_accepted_drafts + 1,
|
||||||
num_accepted_drafts_cpu=num_accepted_drafts_list,
|
num_accepted_drafts_cpu=num_accepted_drafts_list,
|
||||||
@@ -566,7 +566,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
return EagleVerifyOutput(
|
return EagleVerifyOutput(
|
||||||
draft_input=draft_input,
|
draft_input=draft_input,
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
verified_id=verified_id,
|
accept_tokens=accept_tokens,
|
||||||
num_accepted_drafts_per_req_cpu=draft_input.num_accepted_drafts_cpu,
|
num_accepted_drafts_per_req_cpu=draft_input.num_accepted_drafts_cpu,
|
||||||
accepted_indices=accept_index,
|
accepted_indices=accept_index,
|
||||||
)
|
)
|
||||||
@@ -625,7 +625,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
hidden_states=batch.spec_info.hidden_states[
|
hidden_states=batch.spec_info.hidden_states[
|
||||||
unfinished_accept_index
|
unfinished_accept_index
|
||||||
],
|
],
|
||||||
verified_id=predict[unfinished_accept_index],
|
accept_tokens=predict[unfinished_accept_index],
|
||||||
num_accepted_drafts_cpu=draft_input_num_accepted_drafts_cpu,
|
num_accepted_drafts_cpu=draft_input_num_accepted_drafts_cpu,
|
||||||
num_accepted_tokens_cpu=draft_input_num_accepted_tokens_cpu,
|
num_accepted_tokens_cpu=draft_input_num_accepted_tokens_cpu,
|
||||||
num_accepted_drafts=unfinished_num_accepted_drafts,
|
num_accepted_drafts=unfinished_num_accepted_drafts,
|
||||||
@@ -648,7 +648,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
return EagleVerifyOutput(
|
return EagleVerifyOutput(
|
||||||
draft_input=draft_input,
|
draft_input=draft_input,
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
verified_id=verified_id,
|
accept_tokens=accept_tokens,
|
||||||
num_accepted_drafts_per_req_cpu=num_accepted_drafts_list,
|
num_accepted_drafts_per_req_cpu=num_accepted_drafts_list,
|
||||||
accepted_indices=accept_index,
|
accepted_indices=accept_index,
|
||||||
)
|
)
|
||||||
@@ -669,7 +669,14 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
# `num_accepted_drafts` and `num_accepted_tokens` are kept in sync:
|
# `num_accepted_drafts` and `num_accepted_tokens` are kept in sync:
|
||||||
# `num_accepted_tokens = num_accepted_drafts + 1` (per-req, one bonus per req).
|
# `num_accepted_tokens = num_accepted_drafts + 1` (per-req, one bonus per req).
|
||||||
# Storing both avoids repeated `+ 1` at every consumer (attn backends, kernels).
|
# Storing both avoids repeated `+ 1` at every consumer (attn backends, kernels).
|
||||||
verified_id: torch.Tensor = None
|
bonus_tokens: torch.Tensor = None
|
||||||
|
# Flat accepted-token tensor for draft-extend, shape `[sum_accepted]`.
|
||||||
|
# Set right after verify and consumed by `prepare_extend_after_decode` as
|
||||||
|
# the extend batch's `input_ids`. Dead after that method returns.
|
||||||
|
# TODO: drop this field and pass `accept_tokens` directly to
|
||||||
|
# `prepare_extend_after_decode` as a method arg. Its lifetime is bounded
|
||||||
|
# by verify -> prepare_extend, no need to live on the dataclass.
|
||||||
|
accept_tokens: torch.Tensor = None
|
||||||
num_accepted_drafts: torch.Tensor = None
|
num_accepted_drafts: torch.Tensor = None
|
||||||
num_accepted_tokens: torch.Tensor = None
|
num_accepted_tokens: torch.Tensor = None
|
||||||
num_accepted_drafts_cpu: List[int] = None
|
num_accepted_drafts_cpu: List[int] = None
|
||||||
@@ -707,13 +714,13 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
return
|
return
|
||||||
|
|
||||||
# Prefill only generate 1 token.
|
# Prefill only generate 1 token.
|
||||||
assert len(self.verified_id) == len(batch.seq_lens)
|
assert len(self.bonus_tokens) == len(batch.seq_lens)
|
||||||
|
|
||||||
pt = 0
|
pt = 0
|
||||||
for i, extend_len in enumerate(batch.extend_lens):
|
for i, extend_len in enumerate(batch.extend_lens):
|
||||||
input_ids = batch.input_ids[pt : pt + extend_len]
|
input_ids = batch.input_ids[pt : pt + extend_len]
|
||||||
batch.input_ids[pt : pt + extend_len] = torch.cat(
|
batch.input_ids[pt : pt + extend_len] = torch.cat(
|
||||||
(input_ids[1:], self.verified_id[i].reshape(1))
|
(input_ids[1:], self.bonus_tokens[i].reshape(1))
|
||||||
)
|
)
|
||||||
pt += extend_len
|
pt += extend_len
|
||||||
|
|
||||||
@@ -727,7 +734,8 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
capture_hidden_mode: CaptureHiddenMode,
|
capture_hidden_mode: CaptureHiddenMode,
|
||||||
):
|
):
|
||||||
return cls(
|
return cls(
|
||||||
verified_id=torch.empty((0,), device=device, dtype=torch.int32),
|
bonus_tokens=torch.empty((0,), device=device, dtype=torch.int32),
|
||||||
|
accept_tokens=torch.empty((0,), device=device, dtype=torch.int32),
|
||||||
hidden_states=torch.empty((0, hidden_size), device=device, dtype=dtype),
|
hidden_states=torch.empty((0, hidden_size), device=device, dtype=dtype),
|
||||||
topk_p=torch.empty((0, topk), device=device, dtype=torch.float32),
|
topk_p=torch.empty((0, topk), device=device, dtype=torch.float32),
|
||||||
topk_index=torch.empty((0, topk), device=device, dtype=torch.int64),
|
topk_index=torch.empty((0, topk), device=device, dtype=torch.int64),
|
||||||
@@ -748,7 +756,11 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
if batch.forward_mode.is_idle():
|
if batch.forward_mode.is_idle():
|
||||||
return
|
return
|
||||||
|
|
||||||
batch.input_ids = self.verified_id
|
# `self.accept_tokens` is the flat accepted-token tensor set by
|
||||||
|
# `EagleVerifyInput.verify`; use it as the extend batch's `input_ids`.
|
||||||
|
# The kernel below populates `self.bonus_tokens` ([bs] per-req) for
|
||||||
|
# the next decode round.
|
||||||
|
batch.input_ids = self.accept_tokens
|
||||||
batch.extend_lens = batch.spec_info.num_accepted_tokens_cpu
|
batch.extend_lens = batch.spec_info.num_accepted_tokens_cpu
|
||||||
batch.extend_num_tokens = sum(batch.extend_lens)
|
batch.extend_num_tokens = sum(batch.extend_lens)
|
||||||
batch.seq_lens = batch.spec_info.seq_lens_for_draft_extend
|
batch.seq_lens = batch.spec_info.seq_lens_for_draft_extend
|
||||||
@@ -759,14 +771,16 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
|
|
||||||
self.capture_hidden_mode = CaptureHiddenMode.LAST
|
self.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
self.positions = torch.empty_like(batch.input_ids, dtype=torch.long)
|
self.positions = torch.empty_like(batch.input_ids, dtype=torch.long)
|
||||||
self.verified_id = torch.empty_like(self.num_accepted_tokens, dtype=torch.int32)
|
self.bonus_tokens = torch.empty_like(
|
||||||
|
self.num_accepted_tokens, dtype=torch.int32
|
||||||
|
)
|
||||||
|
|
||||||
create_extend_after_decode_spec_info[(len(batch.seq_lens),)](
|
create_extend_after_decode_spec_info[(len(batch.seq_lens),)](
|
||||||
batch.input_ids,
|
batch.input_ids,
|
||||||
batch.seq_lens,
|
batch.seq_lens,
|
||||||
self.num_accepted_tokens,
|
self.num_accepted_tokens,
|
||||||
self.positions,
|
self.positions,
|
||||||
self.verified_id,
|
self.bonus_tokens,
|
||||||
next_power_of_2(max(speculative_num_steps + 1, len(batch.seq_lens))),
|
next_power_of_2(max(speculative_num_steps + 1, len(batch.seq_lens))),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -821,13 +835,13 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
self.topk_p = self.topk_p[: len(new_indices)]
|
self.topk_p = self.topk_p[: len(new_indices)]
|
||||||
self.topk_index = self.topk_index[: len(new_indices)]
|
self.topk_index = self.topk_index[: len(new_indices)]
|
||||||
self.hidden_states = self.hidden_states[: len(new_indices)]
|
self.hidden_states = self.hidden_states[: len(new_indices)]
|
||||||
self.verified_id = self.verified_id[: len(new_indices)]
|
self.bonus_tokens = self.bonus_tokens[: len(new_indices)]
|
||||||
else:
|
else:
|
||||||
# in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index`
|
# in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index`
|
||||||
self.topk_p = self.topk_p[new_indices]
|
self.topk_p = self.topk_p[new_indices]
|
||||||
self.topk_index = self.topk_index[new_indices]
|
self.topk_index = self.topk_index[new_indices]
|
||||||
self.hidden_states = self.hidden_states[new_indices]
|
self.hidden_states = self.hidden_states[new_indices]
|
||||||
self.verified_id = self.verified_id[new_indices]
|
self.bonus_tokens = self.bonus_tokens[new_indices]
|
||||||
|
|
||||||
def merge_batch(self, spec_info: "EagleDraftInput"):
|
def merge_batch(self, spec_info: "EagleDraftInput"):
|
||||||
if self.future_indices is not None:
|
if self.future_indices is not None:
|
||||||
@@ -841,7 +855,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
|
|
||||||
if self.hidden_states is None:
|
if self.hidden_states is None:
|
||||||
self.hidden_states = spec_info.hidden_states
|
self.hidden_states = spec_info.hidden_states
|
||||||
self.verified_id = spec_info.verified_id
|
self.bonus_tokens = spec_info.bonus_tokens
|
||||||
self.topk_p = spec_info.topk_p
|
self.topk_p = spec_info.topk_p
|
||||||
self.topk_index = spec_info.topk_index
|
self.topk_index = spec_info.topk_index
|
||||||
return
|
return
|
||||||
@@ -850,7 +864,9 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
self.hidden_states = torch.cat(
|
self.hidden_states = torch.cat(
|
||||||
[self.hidden_states, spec_info.hidden_states], axis=0
|
[self.hidden_states, spec_info.hidden_states], axis=0
|
||||||
)
|
)
|
||||||
self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0)
|
self.bonus_tokens = torch.cat(
|
||||||
|
[self.bonus_tokens, spec_info.bonus_tokens], axis=0
|
||||||
|
)
|
||||||
self.topk_p = torch.cat([self.topk_p, spec_info.topk_p])
|
self.topk_p = torch.cat([self.topk_p, spec_info.topk_p])
|
||||||
self.topk_index = torch.cat([self.topk_index, spec_info.topk_index])
|
self.topk_index = torch.cat([self.topk_index, spec_info.topk_index])
|
||||||
|
|
||||||
@@ -861,8 +877,8 @@ class EagleVerifyOutput:
|
|||||||
draft_input: EagleDraftInput
|
draft_input: EagleDraftInput
|
||||||
# Logit outputs from target worker
|
# Logit outputs from target worker
|
||||||
logits_output: LogitsProcessorOutput
|
logits_output: LogitsProcessorOutput
|
||||||
# Accepted token ids including the bonus token
|
# Accepted token ids including the bonus token (flat, [sum_accepted])
|
||||||
verified_id: torch.Tensor
|
accept_tokens: torch.Tensor
|
||||||
# Accepted token length per sequence in a batch in CPU.
|
# Accepted token length per sequence in a batch in CPU.
|
||||||
num_accepted_drafts_per_req_cpu: List[int]
|
num_accepted_drafts_per_req_cpu: List[int]
|
||||||
# Accepted indices from logits_output.next_token_logits
|
# Accepted indices from logits_output.next_token_logits
|
||||||
|
|||||||
@@ -473,10 +473,10 @@ class EagleVerifyInputV2Mixin:
|
|||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def fill_new_verified_id(
|
def fill_bonus_tokens(
|
||||||
verified_id,
|
accept_tokens,
|
||||||
accept_lens,
|
accept_lens,
|
||||||
new_verified_id,
|
bonus_tokens_ptr,
|
||||||
num_draft_tokens: tl.constexpr,
|
num_draft_tokens: tl.constexpr,
|
||||||
):
|
):
|
||||||
# NOTE: we cannot fuse any in-place operations of `accept_lens` inside this kernel
|
# NOTE: we cannot fuse any in-place operations of `accept_lens` inside this kernel
|
||||||
@@ -485,9 +485,9 @@ def fill_new_verified_id(
|
|||||||
# `accept_lens` includes the bonus token; the last accepted slot is at -1.
|
# `accept_lens` includes the bonus token; the last accepted slot is at -1.
|
||||||
accept_len = tl.load(accept_lens + pid)
|
accept_len = tl.load(accept_lens + pid)
|
||||||
|
|
||||||
verified_id_idx = num_draft_tokens * pid + accept_len - 1
|
bonus_token_idx = num_draft_tokens * pid + accept_len - 1
|
||||||
verified_id_data = tl.load(verified_id + verified_id_idx)
|
bonus_token = tl.load(accept_tokens + bonus_token_idx)
|
||||||
tl.store(new_verified_id + pid, verified_id_data)
|
tl.store(bonus_tokens_ptr + pid, bonus_token)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
|
|||||||
@@ -46,7 +46,7 @@ class TreeMaskMode(IntEnum):
|
|||||||
|
|
||||||
|
|
||||||
def build_tree_kernel_efficient(
|
def build_tree_kernel_efficient(
|
||||||
verified_id: torch.Tensor,
|
bonus_tokens: torch.Tensor,
|
||||||
parent_list: List[torch.Tensor],
|
parent_list: List[torch.Tensor],
|
||||||
top_scores_index: torch.Tensor,
|
top_scores_index: torch.Tensor,
|
||||||
draft_tokens: torch.Tensor,
|
draft_tokens: torch.Tensor,
|
||||||
@@ -59,7 +59,7 @@ def build_tree_kernel_efficient(
|
|||||||
tree_mask_buf: Optional[torch.Tensor] = None,
|
tree_mask_buf: Optional[torch.Tensor] = None,
|
||||||
position_buf: Optional[torch.Tensor] = None,
|
position_buf: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
draft_tokens = torch.cat((verified_id.unsqueeze(1), draft_tokens), dim=1).flatten()
|
draft_tokens = torch.cat((bonus_tokens.unsqueeze(1), draft_tokens), dim=1).flatten()
|
||||||
|
|
||||||
# seq_lens_sum == sum(seq_lens); seq_lens: sequence length without draft tokens
|
# seq_lens_sum == sum(seq_lens); seq_lens: sequence length without draft tokens
|
||||||
bs = seq_lens.numel()
|
bs = seq_lens.numel()
|
||||||
|
|||||||
@@ -505,7 +505,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
# when DP attention is enabled, but it is slow. Skip it for now.
|
# when DP attention is enabled, but it is slow. Skip it for now.
|
||||||
if (
|
if (
|
||||||
self.server_args.enable_dp_attention
|
self.server_args.enable_dp_attention
|
||||||
or batch.spec_info.verified_id.shape[0] > 0
|
or batch.spec_info.accept_tokens.shape[0] > 0
|
||||||
):
|
):
|
||||||
# decode is not finished
|
# decode is not finished
|
||||||
self.forward_draft_extend_after_decode(batch)
|
self.forward_draft_extend_after_decode(batch)
|
||||||
@@ -521,14 +521,14 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
|
|
||||||
return GenerationBatchResult(
|
return GenerationBatchResult(
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
next_token_ids=verify_output.verified_id,
|
next_token_ids=verify_output.accept_tokens,
|
||||||
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
|
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
|
||||||
num_accepted_drafts_per_req_cpu=verify_output.num_accepted_drafts_per_req_cpu,
|
num_accepted_drafts_per_req_cpu=verify_output.num_accepted_drafts_per_req_cpu,
|
||||||
can_run_cuda_graph=can_run_cuda_graph,
|
can_run_cuda_graph=can_run_cuda_graph,
|
||||||
)
|
)
|
||||||
|
|
||||||
def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch):
|
def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch):
|
||||||
local_need_forward = batch.spec_info.verified_id.shape[0] > 0
|
local_need_forward = batch.spec_info.accept_tokens.shape[0] > 0
|
||||||
if not self.server_args.enable_dp_attention:
|
if not self.server_args.enable_dp_attention:
|
||||||
return local_need_forward
|
return local_need_forward
|
||||||
|
|
||||||
@@ -588,7 +588,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
if batch.sampling_info.penalizer_orchestrator.is_required:
|
if batch.sampling_info.penalizer_orchestrator.is_required:
|
||||||
# This is a relaxed version of penalties for speculative decoding.
|
# This is a relaxed version of penalties for speculative decoding.
|
||||||
batch.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
|
batch.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
|
||||||
spec_info.verified_id.to(torch.int64)
|
spec_info.bonus_tokens.to(torch.int64)
|
||||||
)
|
)
|
||||||
|
|
||||||
# Allocate cache locations
|
# Allocate cache locations
|
||||||
@@ -779,7 +779,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
retrieve_next_sibling,
|
retrieve_next_sibling,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
) = build_tree_kernel_efficient(
|
) = build_tree_kernel_efficient(
|
||||||
spec_info.verified_id,
|
spec_info.bonus_tokens,
|
||||||
parent_list,
|
parent_list,
|
||||||
top_scores_index,
|
top_scores_index,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
@@ -1082,7 +1082,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
"""
|
"""
|
||||||
batch.spec_info = EagleDraftInput(
|
batch.spec_info = EagleDraftInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
verified_id=next_token_ids,
|
bonus_tokens=next_token_ids,
|
||||||
num_tokens_per_req=1,
|
num_tokens_per_req=1,
|
||||||
num_tokens_for_logprob_per_req=1,
|
num_tokens_for_logprob_per_req=1,
|
||||||
)
|
)
|
||||||
@@ -1116,7 +1116,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
|
|
||||||
input_is_idle = batch.forward_mode.is_idle()
|
input_is_idle = batch.forward_mode.is_idle()
|
||||||
|
|
||||||
if not input_is_idle and batch.spec_info.verified_id.numel() == 0:
|
if not input_is_idle and batch.spec_info.accept_tokens.numel() == 0:
|
||||||
batch = batch.copy()
|
batch = batch.copy()
|
||||||
batch.prepare_for_idle()
|
batch.prepare_for_idle()
|
||||||
hidden_size = (
|
hidden_size = (
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
|
|||||||
from sglang.srt.speculative.eagle_info_v2 import (
|
from sglang.srt.speculative.eagle_info_v2 import (
|
||||||
assign_extend_cache_locs,
|
assign_extend_cache_locs,
|
||||||
fill_accepted_out_cache_loc,
|
fill_accepted_out_cache_loc,
|
||||||
fill_new_verified_id,
|
fill_bonus_tokens,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
|
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
@@ -383,7 +383,7 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
retrieve_next_sibling,
|
retrieve_next_sibling,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
) = build_tree_kernel_efficient(
|
) = build_tree_kernel_efficient(
|
||||||
draft_input.verified_id,
|
draft_input.bonus_tokens,
|
||||||
parent_list,
|
parent_list,
|
||||||
top_scores_index,
|
top_scores_index,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
@@ -537,7 +537,7 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
# Construct spec_info
|
# Construct spec_info
|
||||||
next_draft_input = EagleDraftInput(
|
next_draft_input = EagleDraftInput(
|
||||||
hidden_states=target_hidden_states,
|
hidden_states=target_hidden_states,
|
||||||
verified_id=next_token_ids,
|
bonus_tokens=next_token_ids,
|
||||||
new_seq_lens=batch.seq_lens,
|
new_seq_lens=batch.seq_lens,
|
||||||
# draft mode is same with decode mode, only 1 token per req
|
# draft mode is same with decode mode, only 1 token per req
|
||||||
num_tokens_per_req=1,
|
num_tokens_per_req=1,
|
||||||
@@ -1041,16 +1041,16 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
verify_done.record()
|
verify_done.record()
|
||||||
|
|
||||||
if not batch.forward_mode.is_idle():
|
if not batch.forward_mode.is_idle():
|
||||||
all_verified_id = predict[accept_index]
|
accept_tokens = predict[accept_index]
|
||||||
verified_id = torch.empty_like(accept_lens, dtype=torch.int32)
|
bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32)
|
||||||
fill_new_verified_id[(bs,)](
|
fill_bonus_tokens[(bs,)](
|
||||||
all_verified_id,
|
accept_tokens,
|
||||||
accept_lens,
|
accept_lens,
|
||||||
verified_id,
|
bonus_tokens,
|
||||||
self.speculative_num_draft_tokens,
|
self.speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
verified_id = torch.empty((0,), device=self.device, dtype=torch.int32)
|
bonus_tokens = torch.empty((0,), device=self.device, dtype=torch.int32)
|
||||||
|
|
||||||
if batch.return_logprob and not batch.forward_mode.is_idle():
|
if batch.return_logprob and not batch.forward_mode.is_idle():
|
||||||
compute_spec_v2_logprobs(
|
compute_spec_v2_logprobs(
|
||||||
@@ -1059,7 +1059,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
# Construct the next draft input
|
# Construct the next draft input
|
||||||
next_draft_input = EagleDraftInput(
|
next_draft_input = EagleDraftInput(
|
||||||
verified_id=verified_id,
|
bonus_tokens=bonus_tokens,
|
||||||
new_seq_lens=new_seq_lens,
|
new_seq_lens=new_seq_lens,
|
||||||
verify_done=verify_done,
|
verify_done=verify_done,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -137,12 +137,12 @@ def select_last_verified_seed(
|
|||||||
draft_input: FrozenKVMTPDraftInput,
|
draft_input: FrozenKVMTPDraftInput,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
if draft_input.num_accepted_tokens is None:
|
if draft_input.num_accepted_tokens is None:
|
||||||
return draft_input.verified_id, draft_input.hidden_states
|
return draft_input.bonus_tokens, draft_input.hidden_states
|
||||||
|
|
||||||
counts = draft_input.num_accepted_tokens.to(torch.long)
|
counts = draft_input.num_accepted_tokens.to(torch.long)
|
||||||
last_indices = torch.cumsum(counts, dim=0) - 1
|
last_indices = torch.cumsum(counts, dim=0) - 1
|
||||||
return (
|
return (
|
||||||
draft_input.verified_id[last_indices],
|
draft_input.bonus_tokens[last_indices],
|
||||||
draft_input.hidden_states[last_indices],
|
draft_input.hidden_states[last_indices],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -366,7 +366,7 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
if draft_input is None:
|
if draft_input is None:
|
||||||
draft_input = FrozenKVMTPDraftInput()
|
draft_input = FrozenKVMTPDraftInput()
|
||||||
|
|
||||||
draft_input.verified_id = last_token_ids.to(torch.int64)
|
draft_input.bonus_tokens = last_token_ids.to(torch.int64)
|
||||||
draft_input.hidden_states = last_hidden_states
|
draft_input.hidden_states = last_hidden_states
|
||||||
draft_input.capture_hidden_mode = CaptureHiddenMode.LAST
|
draft_input.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
draft_input.num_tokens_per_req = 1
|
draft_input.num_tokens_per_req = 1
|
||||||
@@ -380,7 +380,7 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
spec_info_backup = batch.spec_info
|
spec_info_backup = batch.spec_info
|
||||||
|
|
||||||
batch.forward_mode = ForwardMode.DECODE
|
batch.forward_mode = ForwardMode.DECODE
|
||||||
batch.input_ids = draft_input.verified_id
|
batch.input_ids = draft_input.bonus_tokens
|
||||||
batch.return_hidden_states = False
|
batch.return_hidden_states = False
|
||||||
batch.return_logprob = False
|
batch.return_logprob = False
|
||||||
batch.spec_info = draft_input
|
batch.spec_info = draft_input
|
||||||
@@ -461,14 +461,14 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
if (
|
if (
|
||||||
self.server_args.enable_dp_attention
|
self.server_args.enable_dp_attention
|
||||||
or batch.spec_info.verified_id.numel()
|
or batch.spec_info.bonus_tokens.numel()
|
||||||
):
|
):
|
||||||
self.forward_draft_extend_after_decode(batch)
|
self.forward_draft_extend_after_decode(batch)
|
||||||
set_time_batch(batch.reqs, "set_spec_draft_extend_end_time", trace_only=True)
|
set_time_batch(batch.reqs, "set_spec_draft_extend_end_time", trace_only=True)
|
||||||
|
|
||||||
return GenerationBatchResult(
|
return GenerationBatchResult(
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
next_token_ids=verify_output.verified_id,
|
next_token_ids=verify_output.accept_tokens,
|
||||||
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
|
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
|
||||||
num_accepted_drafts_per_req_cpu=verify_output.num_accepted_drafts_per_req_cpu,
|
num_accepted_drafts_per_req_cpu=verify_output.num_accepted_drafts_per_req_cpu,
|
||||||
can_run_cuda_graph=can_run_cuda_graph,
|
can_run_cuda_graph=can_run_cuda_graph,
|
||||||
@@ -507,7 +507,7 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
def forward_draft_extend_after_decode(self, batch: ScheduleBatch) -> None:
|
def forward_draft_extend_after_decode(self, batch: ScheduleBatch) -> None:
|
||||||
assert isinstance(batch.spec_info, FrozenKVMTPDraftInput)
|
assert isinstance(batch.spec_info, FrozenKVMTPDraftInput)
|
||||||
input_is_idle = batch.forward_mode.is_idle()
|
input_is_idle = batch.forward_mode.is_idle()
|
||||||
if not input_is_idle and batch.spec_info.verified_id.numel() == 0:
|
if not input_is_idle and batch.spec_info.bonus_tokens.numel() == 0:
|
||||||
batch = batch.copy()
|
batch = batch.copy()
|
||||||
batch.prepare_for_idle()
|
batch.prepare_for_idle()
|
||||||
batch.spec_info = FrozenKVMTPDraftInput.create_idle_input(
|
batch.spec_info = FrozenKVMTPDraftInput.create_idle_input(
|
||||||
@@ -564,7 +564,7 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
|
|
||||||
if batch.sampling_info.penalizer_orchestrator.is_required:
|
if batch.sampling_info.penalizer_orchestrator.is_required:
|
||||||
batch.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
|
batch.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
|
||||||
spec_info.verified_id.to(torch.int64)
|
spec_info.bonus_tokens.to(torch.int64)
|
||||||
)
|
)
|
||||||
|
|
||||||
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
|
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
@@ -603,7 +603,7 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
retrieve_next_sibling,
|
retrieve_next_sibling,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
) = build_tree_kernel_efficient(
|
) = build_tree_kernel_efficient(
|
||||||
spec_info.verified_id,
|
spec_info.bonus_tokens,
|
||||||
parent_list,
|
parent_list,
|
||||||
top_scores_index,
|
top_scores_index,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
|
|||||||
@@ -283,20 +283,20 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
# when DP attention is enabled, but it is slow. Skip it for now.
|
# when DP attention is enabled, but it is slow. Skip it for now.
|
||||||
if (
|
if (
|
||||||
self.server_args.enable_dp_attention
|
self.server_args.enable_dp_attention
|
||||||
or batch.spec_info.verified_id.shape[0] > 0
|
or batch.spec_info.accept_tokens.shape[0] > 0
|
||||||
):
|
):
|
||||||
# decode is not finished
|
# decode is not finished
|
||||||
self.forward_draft_extend_after_decode(batch)
|
self.forward_draft_extend_after_decode(batch)
|
||||||
|
|
||||||
return GenerationBatchResult(
|
return GenerationBatchResult(
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
next_token_ids=verify_output.verified_id,
|
next_token_ids=verify_output.accept_tokens,
|
||||||
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
|
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
|
||||||
can_run_cuda_graph=can_run_cuda_graph,
|
can_run_cuda_graph=can_run_cuda_graph,
|
||||||
)
|
)
|
||||||
|
|
||||||
def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch):
|
def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch):
|
||||||
local_need_forward = batch.spec_info.verified_id.shape[0] > 0
|
local_need_forward = batch.spec_info.accept_tokens.shape[0] > 0
|
||||||
if not self.server_args.enable_dp_attention:
|
if not self.server_args.enable_dp_attention:
|
||||||
return local_need_forward
|
return local_need_forward
|
||||||
|
|
||||||
@@ -440,7 +440,7 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
retrieve_next_sibling,
|
retrieve_next_sibling,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
) = build_tree_kernel_efficient(
|
) = build_tree_kernel_efficient(
|
||||||
spec_info.verified_id,
|
spec_info.bonus_tokens,
|
||||||
parent_list,
|
parent_list,
|
||||||
top_scores_index,
|
top_scores_index,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
@@ -609,7 +609,7 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
"""
|
"""
|
||||||
batch.spec_info = EagleDraftInput(
|
batch.spec_info = EagleDraftInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
verified_id=next_token_ids,
|
bonus_tokens=next_token_ids,
|
||||||
num_tokens_per_req=1,
|
num_tokens_per_req=1,
|
||||||
num_tokens_for_logprob_per_req=1,
|
num_tokens_for_logprob_per_req=1,
|
||||||
)
|
)
|
||||||
@@ -664,7 +664,7 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
|
|
||||||
input_is_idle = batch.forward_mode.is_idle()
|
input_is_idle = batch.forward_mode.is_idle()
|
||||||
|
|
||||||
if not input_is_idle and batch.spec_info.verified_id.numel() == 0:
|
if not input_is_idle and batch.spec_info.accept_tokens.numel() == 0:
|
||||||
batch = batch.copy()
|
batch = batch.copy()
|
||||||
batch.prepare_for_idle()
|
batch.prepare_for_idle()
|
||||||
hidden_size = (
|
hidden_size = (
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ from sglang.srt.server_args import ServerArgs
|
|||||||
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
|
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
|
||||||
from sglang.srt.speculative.draft_utils import DraftBackendFactory
|
from sglang.srt.speculative.draft_utils import DraftBackendFactory
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
|
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
|
||||||
from sglang.srt.speculative.eagle_info_v2 import fill_new_verified_id
|
from sglang.srt.speculative.eagle_info_v2 import fill_bonus_tokens
|
||||||
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
|
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
|
||||||
from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
|
from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
|
||||||
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner,
|
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner,
|
||||||
@@ -254,7 +254,7 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
|||||||
retrieve_next_sibling,
|
retrieve_next_sibling,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
) = build_tree_kernel_efficient(
|
) = build_tree_kernel_efficient(
|
||||||
draft_input.verified_id,
|
draft_input.bonus_tokens,
|
||||||
parent_list,
|
parent_list,
|
||||||
top_scores_index,
|
top_scores_index,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
@@ -374,7 +374,7 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
|||||||
# Construct spec_info
|
# Construct spec_info
|
||||||
next_draft_input = EagleDraftInput(
|
next_draft_input = EagleDraftInput(
|
||||||
hidden_states=target_hidden_states,
|
hidden_states=target_hidden_states,
|
||||||
verified_id=next_token_ids,
|
bonus_tokens=next_token_ids,
|
||||||
new_seq_lens=batch.seq_lens,
|
new_seq_lens=batch.seq_lens,
|
||||||
# draft mode is same with decode mode, only 1 token per req
|
# draft mode is same with decode mode, only 1 token per req
|
||||||
num_tokens_per_req=1,
|
num_tokens_per_req=1,
|
||||||
@@ -762,16 +762,16 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
verify_done.record()
|
verify_done.record()
|
||||||
|
|
||||||
if not batch.forward_mode.is_idle():
|
if not batch.forward_mode.is_idle():
|
||||||
all_verified_id = predict[accept_index]
|
accept_tokens = predict[accept_index]
|
||||||
verified_id = torch.empty_like(accept_lens, dtype=torch.int32)
|
bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32)
|
||||||
fill_new_verified_id[(bs,)](
|
fill_bonus_tokens[(bs,)](
|
||||||
all_verified_id,
|
accept_tokens,
|
||||||
accept_lens,
|
accept_lens,
|
||||||
verified_id,
|
bonus_tokens,
|
||||||
self.speculative_num_draft_tokens,
|
self.speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
verified_id = torch.empty((0,), device=self.device, dtype=torch.int32)
|
bonus_tokens = torch.empty((0,), device=self.device, dtype=torch.int32)
|
||||||
|
|
||||||
if batch.return_logprob and not batch.forward_mode.is_idle():
|
if batch.return_logprob and not batch.forward_mode.is_idle():
|
||||||
compute_spec_v2_logprobs(
|
compute_spec_v2_logprobs(
|
||||||
@@ -780,7 +780,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
# Construct the next draft input
|
# Construct the next draft input
|
||||||
next_draft_input = EagleDraftInput(
|
next_draft_input = EagleDraftInput(
|
||||||
verified_id=verified_id,
|
bonus_tokens=bonus_tokens,
|
||||||
new_seq_lens=new_seq_lens,
|
new_seq_lens=new_seq_lens,
|
||||||
verify_done=verify_done,
|
verify_done=verify_done,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -206,7 +206,7 @@ class NgramVerifyInput(SpecInput):
|
|||||||
logits_output.hidden_states = logits_output.hidden_states[
|
logits_output.hidden_states = logits_output.hidden_states[
|
||||||
self.accepted_indices
|
self.accepted_indices
|
||||||
]
|
]
|
||||||
self.verified_id = self.predict[self.accepted_indices]
|
self.accept_tokens = self.predict[self.accepted_indices]
|
||||||
|
|
||||||
def _free_cache(
|
def _free_cache(
|
||||||
self,
|
self,
|
||||||
@@ -465,7 +465,7 @@ class NgramVerifyInput(SpecInput):
|
|||||||
batch.seq_lens.add_(self.num_accepted_tokens)
|
batch.seq_lens.add_(self.num_accepted_tokens)
|
||||||
batch.seq_lens_cpu.add_(num_accepted_tokens_cpu)
|
batch.seq_lens_cpu.add_(num_accepted_tokens_cpu)
|
||||||
|
|
||||||
return logits_output, self.verified_id, num_accepted_drafts
|
return logits_output, self.accept_tokens, num_accepted_drafts
|
||||||
|
|
||||||
def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
|
def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -62,11 +62,11 @@ def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool:
|
|||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def create_extend_after_decode_spec_info(
|
def create_extend_after_decode_spec_info(
|
||||||
verified_id,
|
accept_tokens,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
accept_lens,
|
accept_lens,
|
||||||
positions,
|
positions,
|
||||||
new_verified_id,
|
bonus_tokens_ptr,
|
||||||
bs_upper: tl.constexpr,
|
bs_upper: tl.constexpr,
|
||||||
):
|
):
|
||||||
pid = tl.program_id(axis=0)
|
pid = tl.program_id(axis=0)
|
||||||
@@ -83,8 +83,8 @@ def create_extend_after_decode_spec_info(
|
|||||||
tl.store(positions_ptr + offsets, seq_length - accept_len + offsets, mask)
|
tl.store(positions_ptr + offsets, seq_length - accept_len + offsets, mask)
|
||||||
|
|
||||||
accept_len_cumsum += accept_len - 1
|
accept_len_cumsum += accept_len - 1
|
||||||
verified_id_data = tl.load(verified_id + accept_len_cumsum)
|
bonus_token = tl.load(accept_tokens + accept_len_cumsum)
|
||||||
tl.store(new_verified_id + pid, verified_id_data)
|
tl.store(bonus_tokens_ptr + pid, bonus_token)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ class TestBuildEagleTree(unittest.TestCase):
|
|||||||
|
|
||||||
def test_build_tree_kernel_efficient(self):
|
def test_build_tree_kernel_efficient(self):
|
||||||
"""Test the build_tree_kernel_efficient function with known inputs and expected outputs."""
|
"""Test the build_tree_kernel_efficient function with known inputs and expected outputs."""
|
||||||
verified_id = torch.tensor([29974, 13], device=get_device(), dtype=torch.int32)
|
bonus_tokens = torch.tensor([29974, 13], device=get_device(), dtype=torch.int32)
|
||||||
score_list = [
|
score_list = [
|
||||||
torch.tensor(
|
torch.tensor(
|
||||||
[
|
[
|
||||||
@@ -244,7 +244,7 @@ class TestBuildEagleTree(unittest.TestCase):
|
|||||||
retrieve_next_sibling,
|
retrieve_next_sibling,
|
||||||
draft_tokens,
|
draft_tokens,
|
||||||
) = build_tree_kernel_efficient(
|
) = build_tree_kernel_efficient(
|
||||||
verified_id=verified_id,
|
bonus_tokens=bonus_tokens,
|
||||||
parent_list=parent_list,
|
parent_list=parent_list,
|
||||||
top_scores_index=top_scores_index,
|
top_scores_index=top_scores_index,
|
||||||
draft_tokens=draft_tokens,
|
draft_tokens=draft_tokens,
|
||||||
|
|||||||
Reference in New Issue
Block a user