From f825d729363136a2d4a4b330fa694d0b37a878fa Mon Sep 17 00:00:00 2001 From: Nan Jiang <59716405+nanjiangwill@users.noreply.github.com> Date: Thu, 20 Aug 2026 16:55:25 -0700 Subject: [PATCH] [Sampling] Restore finite top-k requirement for sampling masks (#35205) --- python/sglang/srt/managers/scheduler.py | 10 ++--- .../registered/sampling/test_sampling_mask.py | 40 ++++++++++--------- 2 files changed, 25 insertions(+), 25 deletions(-) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 04dddd34b..58a547b9a 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2572,13 +2572,11 @@ class Scheduler( self._add_request_to_queue(req) return - uses_top_k_or_top_p_truncation = ( - req.sampling_params.top_k != TOP_K_ALL or req.sampling_params.top_p < 1.0 - ) - if req.return_sampling_mask and not uses_top_k_or_top_p_truncation: + if req.return_sampling_mask and req.sampling_params.top_k == TOP_K_ALL: error_msg = ( - "return_sampling_mask cannot return the full vocabulary; set " - "top_p < 1 or a finite top_k." + "return_sampling_mask requires finite top_k; top_p-only sampling " + "is valid but can return huge masks in the tail, blowing up " + "metadata, so we need a safety cap." ) req.set_finish_with_abort(error_msg) self.init_req_max_new_tokens(req) diff --git a/test/registered/sampling/test_sampling_mask.py b/test/registered/sampling/test_sampling_mask.py index 6f04b058b..b86c0d35a 100644 --- a/test/registered/sampling/test_sampling_mask.py +++ b/test/registered/sampling/test_sampling_mask.py @@ -18,13 +18,15 @@ register_amd_ci(est_time=320, suite="stage-b-test-1-gpu-small-amd") _MAX_NEW_TOKENS = 4 _TOP_P = 0.99 -_TOP_P_SMALL = 1e-5 _TOP_K = 10 _SAMPLING_SEED = 1234 _SERVER_ARGS = ( "--mem-fraction-static", "0.7", ) +_INVALID_SAMPLING_MASK_ERROR = ( + "top_p-only sampling is valid but can return huge masks in the tail" +) class SamplingMaskTestMixin: @@ -79,6 +81,11 @@ class SamplingMaskTestMixin: self.assertIn(output_id, sampling_mask) return sampling_masks + def _assert_rejects_unbounded_sampling_mask(self, sampling_params): + response = self._post_generate(sampling_params) + self.assertEqual(response.status_code, 400, response.text) + self.assertIn(_INVALID_SAMPLING_MASK_ERROR, response.text) + class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase): @classmethod @@ -184,24 +191,15 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase): expected_logprob = math.log(probs[output_id] / support_mass) self.assertAlmostEqual(mask_logprob, expected_logprob, delta=1e-2) - def test_generate_returns_top_p_only_sampling_mask(self): - self._generate_sampling_masks( - { - "temperature": 1.0, - "top_p": _TOP_P_SMALL, - "max_new_tokens": _MAX_NEW_TOKENS, - "ignore_eos": True, - } - ) - - def test_chat_completions_returns_top_p_only_sampling_mask(self): + def test_chat_completions_returns_sampling_mask(self): response = requests.post( self.base_url + "/v1/chat/completions", json={ "model": self.model, "messages": [{"role": "user", "content": "Name a capital city."}], "temperature": 1.0, - "top_p": _TOP_P_SMALL, + "top_k": _TOP_K, + "top_p": _TOP_P, "max_tokens": _MAX_NEW_TOKENS, "ignore_eos": True, "return_sampling_mask": True, @@ -224,8 +222,16 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase): for output_id, sampling_mask in zip(output_ids, sampling_masks): self.assertIn(output_id, sampling_mask) - def test_generate_rejects_full_vocabulary_sampling_mask(self): - response = self._post_generate( + def test_generate_rejects_unbounded_sampling_mask(self): + self._assert_rejects_unbounded_sampling_mask( + { + "temperature": 1.0, + "top_p": _TOP_P, + "max_new_tokens": _MAX_NEW_TOKENS, + "ignore_eos": True, + } + ) + self._assert_rejects_unbounded_sampling_mask( { "temperature": 1.0, "top_p": 1.0, @@ -233,10 +239,6 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase): "ignore_eos": True, } ) - self.assertEqual(response.status_code, 400, response.text) - self.assertIn( - "return_sampling_mask cannot return the full vocabulary", response.text - ) class TestSamplingMaskDeterministic(SamplingMaskTestMixin, CustomTestCase):