diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index 85c1526fd..6cd08b92b 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -683,20 +683,22 @@ def eagle_sample( next_token_logits / expanded_temperature, dim=-1 ) # (bs * num_draft_tokens, vocab_size) maybe_detect_nan(target_probs, "v2 verify: target_probs after softmax") - target_probs = top_k_renorm_prob( - target_probs, - torch.repeat_interleave( - sampling_info.top_ks, verify_input.draft_token_num, dim=0 - ), - ) # (bs * num_draft_tokens, vocab_size) - maybe_detect_nan(target_probs, "v2 verify: target_probs after top_k_renorm") - target_probs = top_p_renorm_prob( - target_probs, - torch.repeat_interleave( - sampling_info.top_ps, verify_input.draft_token_num, dim=0 - ), - ) - maybe_detect_nan(target_probs, "v2 verify: target_probs after top_p_renorm") + if sampling_info.need_top_k_sampling: + target_probs = top_k_renorm_prob( + target_probs, + torch.repeat_interleave( + sampling_info.top_ks, verify_input.draft_token_num, dim=0 + ), + ) # (bs * num_draft_tokens, vocab_size) + maybe_detect_nan(target_probs, "v2 verify: target_probs after top_k_renorm") + if sampling_info.need_top_p_sampling: + target_probs = top_p_renorm_prob( + target_probs, + torch.repeat_interleave( + sampling_info.top_ps, verify_input.draft_token_num, dim=0 + ), + ) + maybe_detect_nan(target_probs, "v2 verify: target_probs after top_p_renorm") target_probs = target_probs.reshape(bs, verify_input.draft_token_num, -1) draft_probs = ( verify_input.draft_probs