Simplify prepare_extend_after_decode (#6987)

This commit is contained in:
Lianmin Zheng
2025-06-09 16:39:21 -07:00
committed by GitHub
parent a968c888c0
commit dc0705a504
9 changed files with 140 additions and 176 deletions

View File

@@ -1636,7 +1636,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
if self.spec_info:
self.spec_info.merge_batch(other.spec_info)
def get_model_worker_batch(self) -> ModelWorkerBatch:
def get_model_worker_batch(
self, seq_lens_cpu_cache: Optional[torch.Tensor] = None
) -> ModelWorkerBatch:
if self.forward_mode.is_decode_or_idle():
extend_seq_lens = extend_prefix_lens = extend_logprob_start_lens = None
else:
@@ -1646,16 +1648,20 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# Create seq_lens_cpu when needed
if (
(
global_server_args_dict["attention_backend"] == "fa3"
or (
global_server_args_dict["use_mla_backend"]
and global_server_args_dict["attention_backend"] == "flashinfer"
)
or global_server_args_dict["attention_backend"] == "flashmla"
or global_server_args_dict["attention_backend"] == "fa3"
or global_server_args_dict["attention_backend"] == "cutlass_mla"
or global_server_args_dict["enable_two_batch_overlap"]
):
seq_lens_cpu = self.seq_lens.cpu()
seq_lens_cpu = (
seq_lens_cpu_cache
if seq_lens_cpu_cache is not None
else self.seq_lens.cpu()
)
else:
seq_lens_cpu = None

View File

@@ -1575,10 +1575,9 @@ class Scheduler(
num_accepted_tokens,
can_run_cuda_graph,
) = self.draft_worker.forward_batch_speculative_generation(batch)
self.spec_num_total_accepted_tokens += (
num_accepted_tokens + batch.batch_size()
)
self.spec_num_total_forward_ct += batch.batch_size()
bs = batch.batch_size()
self.spec_num_total_accepted_tokens += num_accepted_tokens + bs
self.spec_num_total_forward_ct += bs
self.num_generated_tokens += num_accepted_tokens
if self.pp_group.is_last_rank: