diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index a4c7c7db0..e947a48cb 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 864acfcc9..7b8c9211b 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -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