Fix overlap scheduler not take effect when outputing logprobs (#14096)

This commit is contained in:
fzyzcjy
2025-11-28 18:15:56 +08:00
committed by GitHub
parent 0e8ce1e832
commit 45cf575852
2 changed files with 6 additions and 6 deletions
+2 -2
View File
@@ -2003,7 +2003,7 @@ class Scheduler(
).Event() ).Event()
if batch_result.delay_sample_func is None: if batch_result.delay_sample_func is None:
self.future_map.store_to_map(future_indices, batch_result) self.future_map.store_to_map(future_indices, batch_result)
batch_result.copy_to_cpu() batch_result.copy_to_cpu(return_logprob=batch.return_logprob)
else: else:
batch_result.future_indices = future_indices batch_result.future_indices = future_indices
@@ -2087,7 +2087,7 @@ class Scheduler(
_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
self.future_map.store_to_map(batch_result.future_indices, batch_result) self.future_map.store_to_map(batch_result.future_indices, batch_result)
batch_result.copy_to_cpu() batch_result.copy_to_cpu(return_logprob=self.cur_batch.return_logprob)
def process_batch_result( def process_batch_result(
self, self,
+4 -4
View File
@@ -43,15 +43,15 @@ class GenerationBatchResult:
# relay path: forward stream -> next step forward # relay path: forward stream -> next step forward
next_draft_input: Optional[EagleDraftInput] = None next_draft_input: Optional[EagleDraftInput] = None
def copy_to_cpu(self, return_logprob: bool = False): def copy_to_cpu(self, return_logprob: bool):
"""Copy tensors to CPU in overlap scheduling. """Copy tensors to CPU in overlap scheduling.
Only the tensors which are needed for processing results are copied, Only the tensors which are needed for processing results are copied,
e.g., next_token_ids, logits outputs e.g., next_token_ids, logits outputs
""" """
if return_logprob: if return_logprob:
if self.logits_output.next_token_logits is not None: if self.logits_output.next_token_logprobs is not None:
self.logits_output.next_token_logits = ( self.logits_output.next_token_logprobs = (
self.logits_output.next_token_logits.to("cpu", non_blocking=True) self.logits_output.next_token_logprobs.to("cpu", non_blocking=True)
) )
if self.logits_output.input_token_logprobs is not None: if self.logits_output.input_token_logprobs is not None:
self.logits_output.input_token_logprobs = ( self.logits_output.input_token_logprobs = (