[Spec] Unify spec/non-spec decode result handling and overlap relay-payload gating (#29225)
This commit is contained in:
@@ -3251,20 +3251,7 @@ class Scheduler(
|
|||||||
# FIXME(lsyin): maybe move this to forward_batch_generation
|
# FIXME(lsyin): maybe move this to forward_batch_generation
|
||||||
batch_result.copy_done = self.device_module.Event()
|
batch_result.copy_done = self.device_module.Event()
|
||||||
if batch_result.delay_sample_func is None:
|
if batch_result.delay_sample_func is None:
|
||||||
# ngram precomputes its draft and does not relay
|
self._relay_forward_payload(future_indices, batch_result)
|
||||||
# through the FutureMap (stash() no-ops for it); its
|
|
||||||
# verify input also has no bonus_tokens to project.
|
|
||||||
if not batch.spec_algorithm.is_ngram():
|
|
||||||
stash_payload = (
|
|
||||||
RelayPayload.from_draft_input(
|
|
||||||
batch_result.next_draft_input
|
|
||||||
)
|
|
||||||
if not batch.spec_algorithm.is_none()
|
|
||||||
else RelayPayload(
|
|
||||||
bonus_tokens=batch_result.next_token_ids
|
|
||||||
)
|
|
||||||
)
|
|
||||||
self.future_map.stash(future_indices, stash_payload)
|
|
||||||
# Result D2H on copy_stream overlaps the next forward
|
# Result D2H on copy_stream overlaps the next forward
|
||||||
# instead of serializing on forward_stream; it's a leaf
|
# instead of serializing on forward_stream; it's a leaf
|
||||||
# gated by copy_done, so nothing on forward_stream waits.
|
# gated by copy_done, so nothing on forward_stream waits.
|
||||||
@@ -3286,11 +3273,7 @@ class Scheduler(
|
|||||||
elif self.enable_pdmux and batch.forward_mode.is_split_prefill():
|
elif self.enable_pdmux and batch.forward_mode.is_split_prefill():
|
||||||
resolve_forward_inputs(batch, self.future_map)
|
resolve_forward_inputs(batch, self.future_map)
|
||||||
batch_result = self.tp_worker.forward_batch_split_prefill(batch)
|
batch_result = self.tp_worker.forward_batch_split_prefill(batch)
|
||||||
if isinstance(batch_result.next_token_ids, torch.Tensor):
|
self._relay_forward_payload(batch.req_pool_indices, batch_result)
|
||||||
self.future_map.stash(
|
|
||||||
batch.req_pool_indices,
|
|
||||||
RelayPayload(bonus_tokens=batch_result.next_token_ids),
|
|
||||||
)
|
|
||||||
batch.input_ids = None
|
batch.input_ids = None
|
||||||
elif not batch.spec_algorithm.is_none():
|
elif not batch.spec_algorithm.is_none():
|
||||||
# Non-overlap: drive the V2 worker synchronously (no
|
# Non-overlap: drive the V2 worker synchronously (no
|
||||||
@@ -3324,12 +3307,9 @@ class Scheduler(
|
|||||||
batch_result = self.model_worker.forward_batch_generation(
|
batch_result = self.model_worker.forward_batch_generation(
|
||||||
batch, **kwargs
|
batch, **kwargs
|
||||||
)
|
)
|
||||||
if isinstance(batch_result.next_token_ids, torch.Tensor):
|
if batch_result.has_sampled_token_ids:
|
||||||
# Non-spec: relay via future_map, gathered next iter.
|
# Non-spec: relay via future_map, gathered next iter.
|
||||||
self.future_map.stash(
|
self._relay_forward_payload(batch.req_pool_indices, batch_result)
|
||||||
batch.req_pool_indices,
|
|
||||||
RelayPayload(bonus_tokens=batch_result.next_token_ids),
|
|
||||||
)
|
|
||||||
batch.input_ids = None
|
batch.input_ids = None
|
||||||
self.update_cache_from_scheduler(batch, batch_result)
|
self.update_cache_from_scheduler(batch, batch_result)
|
||||||
|
|
||||||
@@ -3388,6 +3368,21 @@ class Scheduler(
|
|||||||
ActiveRanksOutput(status=dp_active_ranks.tolist())
|
ActiveRanksOutput(status=dp_active_ranks.tolist())
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _relay_forward_payload(
|
||||||
|
self, future_indices: torch.Tensor, batch_result: GenerationBatchResult
|
||||||
|
) -> None:
|
||||||
|
"""Stash this iter's relay payload for next iter's resolve_forward_inputs.
|
||||||
|
ngram is skipped: it relays its draft via batch.spec_info, not the FutureMap."""
|
||||||
|
if self.spec_algorithm.is_ngram():
|
||||||
|
return
|
||||||
|
if batch_result.next_draft_input is not None:
|
||||||
|
payload = RelayPayload.from_draft_input(batch_result.next_draft_input)
|
||||||
|
elif batch_result.has_sampled_token_ids:
|
||||||
|
payload = RelayPayload(bonus_tokens=batch_result.next_token_ids)
|
||||||
|
else:
|
||||||
|
return
|
||||||
|
self.future_map.stash(future_indices, payload)
|
||||||
|
|
||||||
def launch_batch_sample_if_needed(
|
def launch_batch_sample_if_needed(
|
||||||
self, batch_result: GenerationBatchResult
|
self, batch_result: GenerationBatchResult
|
||||||
) -> Union[GenerationBatchResult]:
|
) -> Union[GenerationBatchResult]:
|
||||||
@@ -3400,11 +3395,8 @@ class Scheduler(
|
|||||||
self.forward_stream.wait_stream(self.schedule_stream)
|
self.forward_stream.wait_stream(self.schedule_stream)
|
||||||
_batch_result = batch_result.delay_sample_func()
|
_batch_result = batch_result.delay_sample_func()
|
||||||
assert _batch_result is batch_result
|
assert _batch_result is batch_result
|
||||||
# Delay-sample is non-spec only; stash takes next_token_ids tensor.
|
# Delay-sample is non-spec only; relays the sampled bonus tokens.
|
||||||
self.future_map.stash(
|
self._relay_forward_payload(batch_result.future_indices, batch_result)
|
||||||
batch_result.future_indices,
|
|
||||||
RelayPayload(bonus_tokens=batch_result.next_token_ids),
|
|
||||||
)
|
|
||||||
batch_result.copy_to_cpu(
|
batch_result.copy_to_cpu(
|
||||||
return_logprob=self.cur_batch.return_logprob,
|
return_logprob=self.cur_batch.return_logprob,
|
||||||
return_hidden_states=self.cur_batch.return_hidden_states,
|
return_hidden_states=self.cur_batch.return_hidden_states,
|
||||||
|
|||||||
@@ -675,18 +675,11 @@ class SchedulerBatchResultProcessor:
|
|||||||
# And all the over-allocated tokens will be freed in `release_kv_cache`.
|
# And all the over-allocated tokens will be freed in `release_kv_cache`.
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Non-spec and Spec V2: full post-processing.
|
# next_token_id is a per-req list: 1 token for non-spec, the verified
|
||||||
|
# run for spec (already grammar-truncated in _resolve_spec_v2_tokens).
|
||||||
next_token_id = next_token_ids[i]
|
next_token_id = next_token_ids[i]
|
||||||
is_spec = not batch.spec_algorithm.is_none()
|
is_spec = not batch.spec_algorithm.is_none()
|
||||||
|
|
||||||
if not is_spec:
|
|
||||||
# Normal decode: a single sampled token.
|
|
||||||
req.output_ids.append(next_token_id)
|
|
||||||
new_accept_len = 1
|
|
||||||
else:
|
|
||||||
# Spec: accept the whole verified run. For grammar requests the
|
|
||||||
# run was already truncated at the grammar-terminating token in
|
|
||||||
# _resolve_spec_v2_tokens, so nothing is emitted past completion.
|
|
||||||
req.output_ids.extend(next_token_id)
|
req.output_ids.extend(next_token_id)
|
||||||
new_accept_len = len(next_token_id)
|
new_accept_len = len(next_token_id)
|
||||||
|
|
||||||
@@ -707,21 +700,14 @@ class SchedulerBatchResultProcessor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if req.return_hidden_states and logits_output.hidden_states is not None:
|
if req.return_hidden_states and logits_output.hidden_states is not None:
|
||||||
if not is_spec:
|
# hidden_states is [bs * stride, hidden_dim], one row per emitted
|
||||||
req.hidden_states.append(
|
# token; stride = speculative_num_draft_tokens for spec, 1 for non-spec.
|
||||||
logits_output.hidden_states[i].cpu().clone().tolist()
|
stride = result.speculative_num_draft_tokens or 1
|
||||||
)
|
|
||||||
else:
|
|
||||||
# Spec V2: hidden_states is [bs * speculative_num_draft_tokens, hidden_dim].
|
|
||||||
# One row per emitted token; next_token_id is already truncated
|
|
||||||
# at grammar termination, so this stays aligned with output_ids.
|
|
||||||
stride = result.speculative_num_draft_tokens
|
|
||||||
accept_len = len(next_token_id)
|
accept_len = len(next_token_id)
|
||||||
start = i * stride
|
start = i * stride
|
||||||
req.hidden_states.extend(
|
req.hidden_states.extend(
|
||||||
logits_output.hidden_states[start : start + accept_len]
|
logits_output.hidden_states[start : start + accept_len]
|
||||||
.cpu()
|
.cpu()
|
||||||
.clone()
|
|
||||||
.tolist()
|
.tolist()
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -753,12 +739,18 @@ class SchedulerBatchResultProcessor:
|
|||||||
next_token_ids: Union[torch.Tensor, List[int]],
|
next_token_ids: Union[torch.Tensor, List[int]],
|
||||||
) -> Tuple[Union[List[int], List[List[int]]], Optional[List[float]]]:
|
) -> Tuple[Union[List[int], List[List[int]]], Optional[List[float]]]:
|
||||||
next_token_logprobs = None
|
next_token_logprobs = None
|
||||||
|
# Normalize to a uniform per-req list of accepted tokens (List[List[int]]):
|
||||||
|
# spec unpacks the padded verify output; non-spec wraps its single token.
|
||||||
if not batch.spec_algorithm.is_none():
|
if not batch.spec_algorithm.is_none():
|
||||||
next_token_ids = self._resolve_spec_v2_tokens(result, batch)
|
next_token_ids = self._resolve_spec_v2_tokens(result, batch)
|
||||||
elif isinstance(next_token_ids, list):
|
|
||||||
pass # MLX path: already a list[int], skip torch round-trip
|
|
||||||
else:
|
else:
|
||||||
next_token_ids = next_token_ids.tolist()
|
# CUDA workers return a device tensor, MLX a host list[int]; both -> list.
|
||||||
|
ids = (
|
||||||
|
next_token_ids.tolist()
|
||||||
|
if torch.is_tensor(next_token_ids)
|
||||||
|
else next_token_ids
|
||||||
|
)
|
||||||
|
next_token_ids = [[t] for t in ids]
|
||||||
|
|
||||||
if batch.return_logprob:
|
if batch.return_logprob:
|
||||||
next_token_logprobs = logits_output.next_token_logprobs.tolist()
|
next_token_logprobs = logits_output.next_token_logprobs.tolist()
|
||||||
@@ -786,14 +778,15 @@ class SchedulerBatchResultProcessor:
|
|||||||
next_token_logprobs: list,
|
next_token_logprobs: list,
|
||||||
logits_output: LogitsProcessorOutput,
|
logits_output: LogitsProcessorOutput,
|
||||||
) -> None:
|
) -> None:
|
||||||
# Normalize: non-spec has 1 token, spec decoding has multiple.
|
# accepted_ids is already a per-req list; non-spec logprobs are flat, so
|
||||||
|
# the scalar logprob still needs wrapping.
|
||||||
if not batch.spec_algorithm.is_none():
|
if not batch.spec_algorithm.is_none():
|
||||||
accepted_logprobs = next_token_logprobs[i]
|
accepted_logprobs = next_token_logprobs[i]
|
||||||
accepted_ids = next_token_id
|
accepted_ids = next_token_id
|
||||||
max_accept = len(accepted_logprobs)
|
max_accept = len(accepted_logprobs)
|
||||||
else:
|
else:
|
||||||
accepted_logprobs = [next_token_logprobs[i]]
|
accepted_logprobs = [next_token_logprobs[i]]
|
||||||
accepted_ids = [next_token_id]
|
accepted_ids = next_token_id
|
||||||
max_accept = 1
|
max_accept = 1
|
||||||
|
|
||||||
for j, tok_id in enumerate(accepted_ids):
|
for j, tok_id in enumerate(accepted_ids):
|
||||||
|
|||||||
@@ -85,6 +85,12 @@ class GenerationBatchResult:
|
|||||||
fpm_start_event: Optional[torch.cuda.Event] = None
|
fpm_start_event: Optional[torch.cuda.Event] = None
|
||||||
fpm_end_event: Optional[torch.cuda.Event] = None
|
fpm_end_event: Optional[torch.cuda.Event] = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def has_sampled_token_ids(self) -> bool:
|
||||||
|
"""True when this iter sampled token ids; False when none were produced
|
||||||
|
this rank/split (a non-last PP rank or a non-final prefill split)."""
|
||||||
|
return isinstance(self.next_token_ids, torch.Tensor)
|
||||||
|
|
||||||
@torch.profiler.record_function("copy_result_to_cpu")
|
@torch.profiler.record_function("copy_result_to_cpu")
|
||||||
def copy_to_cpu(self, return_logprob: bool, return_hidden_states: bool = True):
|
def copy_to_cpu(self, return_logprob: bool, return_hidden_states: bool = True):
|
||||||
"""Copy tensors to CPU in overlap scheduling.
|
"""Copy tensors to CPU in overlap scheduling.
|
||||||
|
|||||||
Reference in New Issue
Block a user