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_val,
|
||||||
logits_output.next_token_top_logprobs_idx,
|
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):
|
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_val,
|
||||||
logits_output.next_token_token_ids_logprobs_idx,
|
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[
|
logits_output.next_token_logprobs = logprobs[
|
||||||
torch.arange(len(batch_next_token_ids), device=sampling_info.device),
|
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_val,
|
||||||
logits_output.next_token_top_logprobs_idx,
|
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
|
# Handle token_ids logprobs if requested
|
||||||
if needs_token_ids_logprobs:
|
if needs_token_ids_logprobs:
|
||||||
|
|||||||
@@ -152,6 +152,18 @@ class SchedulerOutputProcessorMixin:
|
|||||||
logits_output.input_token_logprobs = tuple(
|
logits_output.input_token_logprobs = tuple(
|
||||||
logits_output.input_token_logprobs.tolist()
|
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
|
hidden_state_offset = 0
|
||||||
|
|
||||||
@@ -377,7 +389,7 @@ class SchedulerOutputProcessorMixin:
|
|||||||
|
|
||||||
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()
|
||||||
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 = [
|
logits_output.next_token_top_logprobs_val = [
|
||||||
v.tolist() for v in 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
|
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 = [
|
logits_output.next_token_token_ids_logprobs_val = [
|
||||||
v.tolist()
|
v.tolist()
|
||||||
for v in logits_output.next_token_token_ids_logprobs_val
|
for v in logits_output.next_token_token_ids_logprobs_val
|
||||||
|
|||||||
Reference in New Issue
Block a user