Remove sync when enabling return_logprob (#20972)
This commit is contained in:
@@ -327,13 +327,15 @@ class Sampler(nn.Module):
|
||||
(
|
||||
logits_output.next_token_top_logprobs_val,
|
||||
logits_output.next_token_top_logprobs_idx,
|
||||
) = get_top_logprobs(logprobs, top_logprobs_nums)
|
||||
) = get_top_logprobs(logprobs, top_logprobs_nums, no_copy_to_cpu=True)
|
||||
|
||||
if any(x is not None for x in token_ids_logprobs):
|
||||
(
|
||||
logits_output.next_token_token_ids_logprobs_val,
|
||||
logits_output.next_token_token_ids_logprobs_idx,
|
||||
) = get_token_ids_logprobs(logprobs, token_ids_logprobs)
|
||||
) = get_token_ids_logprobs(
|
||||
logprobs, token_ids_logprobs, no_copy_to_cpu=True
|
||||
)
|
||||
|
||||
logits_output.next_token_logprobs = logprobs[
|
||||
torch.arange(len(batch_next_token_ids), device=sampling_info.device),
|
||||
@@ -397,7 +399,7 @@ class Sampler(nn.Module):
|
||||
(
|
||||
logits_output.next_token_top_logprobs_val,
|
||||
logits_output.next_token_top_logprobs_idx,
|
||||
) = get_top_logprobs(logprobs, top_logprobs_nums)
|
||||
) = get_top_logprobs(logprobs, top_logprobs_nums, no_copy_to_cpu=True)
|
||||
|
||||
# Handle token_ids logprobs if requested
|
||||
if needs_token_ids_logprobs:
|
||||
|
||||
@@ -152,6 +152,18 @@ class SchedulerOutputProcessorMixin:
|
||||
logits_output.input_token_logprobs = tuple(
|
||||
logits_output.input_token_logprobs.tolist()
|
||||
)
|
||||
if logits_output.next_token_top_logprobs_val:
|
||||
logits_output.next_token_top_logprobs_val = [
|
||||
v.tolist() for v in logits_output.next_token_top_logprobs_val
|
||||
]
|
||||
logits_output.next_token_top_logprobs_idx = [
|
||||
x.tolist() for x in logits_output.next_token_top_logprobs_idx
|
||||
]
|
||||
if logits_output.next_token_token_ids_logprobs_val:
|
||||
logits_output.next_token_token_ids_logprobs_val = [
|
||||
v.tolist()
|
||||
for v in logits_output.next_token_token_ids_logprobs_val
|
||||
]
|
||||
|
||||
hidden_state_offset = 0
|
||||
|
||||
@@ -377,7 +389,7 @@ class SchedulerOutputProcessorMixin:
|
||||
|
||||
if batch.return_logprob:
|
||||
next_token_logprobs = logits_output.next_token_logprobs.tolist()
|
||||
if batch.is_spec_v2 and logits_output.next_token_top_logprobs_val:
|
||||
if logits_output.next_token_top_logprobs_val:
|
||||
logits_output.next_token_top_logprobs_val = [
|
||||
v.tolist() for v in logits_output.next_token_top_logprobs_val
|
||||
]
|
||||
@@ -385,7 +397,7 @@ class SchedulerOutputProcessorMixin:
|
||||
x.tolist() for x in logits_output.next_token_top_logprobs_idx
|
||||
]
|
||||
|
||||
if batch.is_spec_v2 and logits_output.next_token_token_ids_logprobs_val:
|
||||
if logits_output.next_token_token_ids_logprobs_val:
|
||||
logits_output.next_token_token_ids_logprobs_val = [
|
||||
v.tolist()
|
||||
for v in logits_output.next_token_token_ids_logprobs_val
|
||||
|
||||
Reference in New Issue
Block a user