drop output ids (#25774)
This commit is contained in:
@@ -467,7 +467,7 @@ def extend(reqs, model_runner):
|
||||
|
||||
@torch.no_grad
|
||||
def decode(input_token_ids, batch, model_runner):
|
||||
batch.output_ids = input_token_ids
|
||||
batch.input_ids = input_token_ids.to(torch.int64)
|
||||
batch.prepare_for_decode()
|
||||
_maybe_prepare_mlp_sync_batch(batch, model_runner)
|
||||
forward_batch = ForwardBatch.init_new(batch, model_runner)
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from http import HTTPStatus
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
import torch
|
||||
|
||||
@@ -107,9 +107,9 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
||||
future_map: FutureMap,
|
||||
):
|
||||
"""Assign the buffered last input id to schedule batch"""
|
||||
self.output_ids = []
|
||||
last_tokens: List[int] = []
|
||||
for req in self.reqs:
|
||||
self.output_ids.append(req.output_ids[-1])
|
||||
last_tokens.append(req.output_ids[-1])
|
||||
maybe_cache_unfinished_req(req, self.tree_cache)
|
||||
if req.grammar is not None:
|
||||
# FIXME: this try-except block is for handling unexpected xgrammar issue.
|
||||
@@ -130,7 +130,9 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
||||
error_message, HTTPStatus.INTERNAL_SERVER_ERROR
|
||||
)
|
||||
req.grammar.finished = req.finished()
|
||||
self.output_ids = torch.tensor(self.output_ids, device=self.device)
|
||||
last_tokens_tensor = torch.tensor(
|
||||
last_tokens, dtype=torch.int64, device=self.device
|
||||
)
|
||||
|
||||
# Simulate the eagle run.
|
||||
if self.spec_algorithm.is_eagle():
|
||||
@@ -170,9 +172,11 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
||||
topk_p=topk_p,
|
||||
topk_index=topk_index,
|
||||
hidden_states=hidden_states,
|
||||
bonus_tokens=self.output_ids,
|
||||
bonus_tokens=last_tokens_tensor,
|
||||
new_seq_lens=self.seq_lens,
|
||||
)
|
||||
# prepare_for_extend shifts batch.input_ids in place — keep it
|
||||
# as the prefill prompt, not the [bs] last-token tensor.
|
||||
spec_info.prepare_for_extend(self)
|
||||
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||
if self.enable_overlap:
|
||||
@@ -183,3 +187,6 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
||||
spec_info.future_indices, spec_info
|
||||
)
|
||||
self.spec_info = spec_info
|
||||
else:
|
||||
# Non-spec: input_ids feeds the next decode forward directly.
|
||||
self.input_ids = last_tokens_tensor
|
||||
|
||||
@@ -132,7 +132,7 @@ class SchedulerMlxOverlapMixin:
|
||||
pending.reqs,
|
||||
)
|
||||
if result.next_token_ids is not None:
|
||||
pending.batch_copy.output_ids = result.next_token_ids
|
||||
pending.batch_copy.input_ids = result.next_token_ids
|
||||
self.process_batch_result(pending.batch_copy, result)
|
||||
|
||||
def _launch_fresh(batch: "ScheduleBatch") -> MlxPendingJob:
|
||||
|
||||
@@ -1463,7 +1463,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
seq_lens_cpu: torch.Tensor = None # shape: [b], int64
|
||||
# The output locations of the KV cache
|
||||
out_cache_loc: torch.Tensor = None # shape: [b], int64
|
||||
output_ids: torch.Tensor = None # shape: [b], int64
|
||||
|
||||
# For hybrid GDN prefix cache
|
||||
mamba_track_indices: torch.Tensor = None # shape: [b], int64
|
||||
@@ -2392,15 +2391,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
)
|
||||
else:
|
||||
self.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
|
||||
self.output_ids.to(torch.int64)
|
||||
self.input_ids
|
||||
)
|
||||
|
||||
# Update fields
|
||||
# Coerce to int64: torch sampling helpers (sampling_from_probs_torch /
|
||||
# top_k_top_p_min_p_sampling_from_probs_torch) return int32 token ids,
|
||||
# but downstream kernels enforce int64 (e.g. DeepSeek-V4 hash_topk).
|
||||
self.input_ids = self.output_ids.to(torch.int64)
|
||||
self.output_ids = None
|
||||
# input_ids is set at end of previous run_batch (placeholder for
|
||||
# overlap; next_token_ids cast for non-overlap).
|
||||
|
||||
if self.model_config.is_encoder_decoder:
|
||||
self.prepare_encoder_info_decode()
|
||||
@@ -2512,8 +2507,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self.out_cache_loc = None
|
||||
self.seq_lens_sum = self.seq_lens.sum().item()
|
||||
|
||||
if self.output_ids is not None:
|
||||
self.output_ids = self.output_ids[keep_indices_device]
|
||||
if self.input_ids is not None:
|
||||
self.input_ids = self.input_ids[keep_indices_device]
|
||||
|
||||
self.mamba_track_indices = None
|
||||
self.mamba_track_mask = None
|
||||
@@ -2571,8 +2566,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self.orig_seq_lens = torch.cat([self.orig_seq_lens, other.orig_seq_lens])
|
||||
self.out_cache_loc = None
|
||||
self.seq_lens_sum += other.seq_lens_sum
|
||||
if self.output_ids is not None:
|
||||
self.output_ids = torch.cat([self.output_ids, other.output_ids])
|
||||
if self.input_ids is not None:
|
||||
self.input_ids = torch.cat([self.input_ids, other.input_ids])
|
||||
self.mamba_track_indices = None
|
||||
self.mamba_track_mask = None
|
||||
self.mamba_track_seqlens = None
|
||||
|
||||
@@ -2245,7 +2245,7 @@ class Scheduler(
|
||||
batch.seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int64)
|
||||
batch.orig_seq_lens = torch.tensor(seq_lens, dtype=torch.int32, device=device)
|
||||
batch.seq_lens_sum = sum(seq_lens)
|
||||
batch.output_ids = torch.tensor(
|
||||
batch.input_ids = torch.tensor(
|
||||
[r.output_ids[-1] for r in reqs], dtype=torch.int64, device=device
|
||||
)
|
||||
|
||||
@@ -2862,8 +2862,9 @@ class Scheduler(
|
||||
else:
|
||||
batch_result.future_indices = future_indices
|
||||
|
||||
# FIXME(lsyin): move this assignment elsewhere
|
||||
future_indices_or_next_token_ids = -future_indices.indices
|
||||
# Placeholder for next iter's resolve_future to look up the
|
||||
# real token from token_ids_buf via the negated indices.
|
||||
batch.input_ids = -future_indices.indices
|
||||
|
||||
if batch.is_spec_v2:
|
||||
# FIXME(lsyin): tmp code for spec v2
|
||||
@@ -2877,7 +2878,8 @@ class Scheduler(
|
||||
batch.seq_lens = batch_result.next_draft_input.new_seq_lens
|
||||
elif self.enable_pdmux and batch.forward_mode.is_split_prefill():
|
||||
batch_result = self.tp_worker.forward_batch_split_prefill(batch)
|
||||
future_indices_or_next_token_ids = batch_result.next_token_ids
|
||||
if isinstance(batch_result.next_token_ids, torch.Tensor):
|
||||
batch.input_ids = batch_result.next_token_ids.to(torch.int64)
|
||||
else:
|
||||
kwargs = (
|
||||
{"pp_proxy_tensors": pp_proxy_tensors}
|
||||
@@ -2887,15 +2889,12 @@ class Scheduler(
|
||||
batch_result = self.model_worker.forward_batch_generation(
|
||||
batch, **kwargs
|
||||
)
|
||||
future_indices_or_next_token_ids = batch_result.next_token_ids
|
||||
# PP intermediate ranks return None; DLLM returns a per-req list.
|
||||
# Only the tensor case maps onto batch.input_ids as next-iter input.
|
||||
if isinstance(batch_result.next_token_ids, torch.Tensor):
|
||||
batch.input_ids = batch_result.next_token_ids.to(torch.int64)
|
||||
self.update_cache_from_scheduler(batch, batch_result)
|
||||
|
||||
# NOTE: future_indices_or_next_token_ids is used in ScheduleBatch,
|
||||
# which can probably be replaced by future_indices later [TODO(lsyin)].
|
||||
# we shall still keep the original outputs, e.g. next_token_ids
|
||||
# in the GenerationBatchOutput for processing after copy_done.
|
||||
batch.output_ids = future_indices_or_next_token_ids
|
||||
|
||||
# These 2 values are needed for processing the output, but the values can be
|
||||
# modified by overlap schedule. So we have to copy them here so that
|
||||
# we can use the correct values in output processing.
|
||||
|
||||
@@ -1046,7 +1046,7 @@ class SchedulerPPMixin:
|
||||
extend_input_len_per_req,
|
||||
extend_logprob_start_len_per_req,
|
||||
) = get_logprob_from_pp_outputs(pp_outputs)
|
||||
batch.output_ids = pp_outputs["next_token_ids"]
|
||||
batch.input_ids = pp_outputs["next_token_ids"].to(torch.int64)
|
||||
output_result = GenerationBatchResult(
|
||||
logits_output=logits_output,
|
||||
pp_hidden_states_proxy_tensors=None,
|
||||
|
||||
Reference in New Issue
Block a user