Fix spec decoding with grammar in disagg (#24082)
Co-authored-by: jimmy.shong <jimmy.shong@radixark.ai> Co-authored-by: Jimmy Shong <69131491+Jiminator@users.noreply.github.com> Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com> Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> Co-authored-by: Codex <codex@example.com>
This commit is contained in:
co-authored by
jimmy.shong
Jimmy Shong
Xinyuan Tong
Xinyuan Tong
Codex
parent
bea282cede
commit
9fc9d37f6d
@@ -648,16 +648,25 @@ class SchedulerBatchResultProcessor:
|
||||
# Non-spec and V2: full post-processing
|
||||
next_token_id = next_token_ids[i]
|
||||
new_accepted_len = 1
|
||||
needs_tokenwise_grammar_accept = (
|
||||
not batch.spec_algorithm.is_none() and req.grammar is not None
|
||||
)
|
||||
if batch.spec_algorithm.is_none():
|
||||
req.output_ids.append(next_token_id)
|
||||
elif needs_tokenwise_grammar_accept:
|
||||
next_token_id = self._accept_spec_v2_grammar_tokens(req, next_token_id)
|
||||
next_token_ids[i] = next_token_id
|
||||
new_accepted_len = len(next_token_id)
|
||||
else:
|
||||
req.output_ids.extend(next_token_id)
|
||||
new_accepted_len = len(next_token_id)
|
||||
|
||||
self._maybe_update_reasoning_tokens(req, next_token_id)
|
||||
if not needs_tokenwise_grammar_accept:
|
||||
self._maybe_update_reasoning_tokens(req, next_token_id)
|
||||
|
||||
req.time_stats.set_last_decode_finish_time()
|
||||
req.update_finish_state(new_accepted_len)
|
||||
if not needs_tokenwise_grammar_accept:
|
||||
req.update_finish_state(new_accepted_len)
|
||||
|
||||
self._handle_finish_state_updated_req(req, batch, result, i, logits_output)
|
||||
|
||||
@@ -689,9 +698,13 @@ class SchedulerBatchResultProcessor:
|
||||
)
|
||||
|
||||
if req.grammar is not None:
|
||||
self._apply_decode_grammar(
|
||||
req=req, next_token_id=next_token_id, batch=batch
|
||||
)
|
||||
if needs_tokenwise_grammar_accept:
|
||||
# Already advanced token-by-token above; just sync terminal flag.
|
||||
req.grammar.finished = req.finished()
|
||||
else:
|
||||
self._apply_decode_grammar(
|
||||
req=req, next_token_id=next_token_id, batch=batch
|
||||
)
|
||||
|
||||
self.output_streamer.stream_output(batch.reqs, batch.return_logprob)
|
||||
self.token_to_kv_pool_allocator.free_group_end()
|
||||
@@ -777,6 +790,40 @@ class SchedulerBatchResultProcessor:
|
||||
logits_output.next_token_token_ids_logprobs_idx[flat_idx]
|
||||
)
|
||||
|
||||
def _accept_spec_v2_grammar_tokens(
|
||||
self, req: Req, proposed: List[int]
|
||||
) -> List[int]:
|
||||
"""Accept speculative grammar tokens until the request finishes.
|
||||
|
||||
Returns the retained prefix and rolls back KV commits for dropped suffix
|
||||
tokens.
|
||||
"""
|
||||
accept_tokens = []
|
||||
try:
|
||||
for token_id in proposed:
|
||||
req.grammar.accept_token(token_id)
|
||||
req.output_ids.append(token_id)
|
||||
accept_tokens.append(token_id)
|
||||
self._maybe_update_reasoning_tokens(req, token_id)
|
||||
req.update_finish_state()
|
||||
if req.finished():
|
||||
break
|
||||
except ValueError as e:
|
||||
# accept_token raises ValueError if the token is not in the grammar
|
||||
# (misconfigured grammar or invalid token); abort the request.
|
||||
logger.error(
|
||||
f"Grammar accept_token failed for req {req.rid} with token {proposed}: {e}"
|
||||
)
|
||||
self.abort_request(AbortReq(rid=req.rid))
|
||||
req.update_finish_state()
|
||||
|
||||
# _resolve_spec_v2_tokens committed the full proposed list; rollback the
|
||||
# suffix that grammar termination dropped.
|
||||
dropped = len(proposed) - len(accept_tokens)
|
||||
if dropped > 0:
|
||||
req.kv_committed_len -= dropped
|
||||
return accept_tokens
|
||||
|
||||
def _apply_decode_grammar(
|
||||
self,
|
||||
*,
|
||||
|
||||
Reference in New Issue
Block a user