From be804c1b83870947f7c0141ada8646975131b0e8 Mon Sep 17 00:00:00 2001 From: Nan Jiang <59716405+nanjiangwill@users.noreply.github.com> Date: Fri, 14 Aug 2026 16:40:27 -0700 Subject: [PATCH] [RL] Expose top-p-only sampling masks (#33593) --- .../sglang/srt/entrypoints/openai/protocol.py | 1 + .../srt/entrypoints/openai/serving_chat.py | 4 ++ python/sglang/srt/managers/scheduler.py | 10 ++-- .../registered/sampling/test_sampling_mask.py | 53 ++++++++++++++----- .../unit/entrypoints/openai/test_protocol.py | 1 + .../entrypoints/openai/test_serving_chat.py | 15 ++++++ 6 files changed, 68 insertions(+), 16 deletions(-) diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index 5652d63fd..e473976a8 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -795,6 +795,7 @@ class ChatCompletionRequest(BaseModel): return_prompt_token_ids: bool = False return_token_ids: bool = False return_meta_info: bool = False + return_sampling_mask: bool = False reasoning_effort: ReasoningEffortType = Field( default=None, description="Constrains effort on reasoning for reasoning models. " diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index dd0e3ab55..09f7ff06c 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -816,6 +816,9 @@ class OpenAIServingChat(OpenAIServingBase): if not request.messages: return "Messages cannot be empty." + if request.return_sampling_mask and not request.return_meta_info: + return "return_sampling_mask requires return_meta_info=true." + media_error = self._validate_media_content(request) if media_error: return media_error @@ -1004,6 +1007,7 @@ class OpenAIServingChat(OpenAIServingBase): return_logprob=request.logprobs, logprob_start_len=-1, top_logprobs_num=request.top_logprobs or 0, + return_sampling_mask=request.return_sampling_mask, stream=request.stream, return_text_in_logprobs=True, modalities=processed_messages.modalities, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index b8b0134d8..3b66940bd 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2528,11 +2528,13 @@ class Scheduler( self._add_request_to_queue(req) return - if req.return_sampling_mask and req.sampling_params.top_k == TOP_K_ALL: + 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: error_msg = ( - "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." + "return_sampling_mask cannot return the full vocabulary; set " + "top_p < 1 or a finite top_k." ) 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 dbedf2cb0..b68a08ea6 100644 --- a/test/registered/sampling/test_sampling_mask.py +++ b/test/registered/sampling/test_sampling_mask.py @@ -18,15 +18,13 @@ 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: @@ -81,11 +79,6 @@ 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 @@ -191,16 +184,48 @@ 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_rejects_unbounded_sampling_mask(self): - self._assert_rejects_unbounded_sampling_mask( + def test_generate_returns_top_p_only_sampling_mask(self): + self._generate_sampling_masks( { "temperature": 1.0, - "top_p": _TOP_P, + "top_p": _TOP_P_SMALL, "max_new_tokens": _MAX_NEW_TOKENS, "ignore_eos": True, } ) - self._assert_rejects_unbounded_sampling_mask( + + def test_chat_completions_returns_top_p_only_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, + "max_tokens": _MAX_NEW_TOKENS, + "ignore_eos": True, + "return_sampling_mask": True, + "return_meta_info": True, + "return_token_ids": True, + }, + timeout=60, + ) + self.assertEqual(response.status_code, 200, response.text) + + choice = response.json()["choices"][0] + output_ids = choice["token_ids"] + meta_info = choice["meta_info"] + sampling_masks = meta_info["output_token_sampling_mask"] + sampling_logprobs = meta_info["output_token_sampling_logprobs"] + + self.assertEqual(len(output_ids), _MAX_NEW_TOKENS) + self.assertEqual(len(sampling_masks), len(output_ids)) + self.assertEqual(len(sampling_logprobs), len(output_ids)) + 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( { "temperature": 1.0, "top_p": 1.0, @@ -208,6 +233,10 @@ 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): diff --git a/test/registered/unit/entrypoints/openai/test_protocol.py b/test/registered/unit/entrypoints/openai/test_protocol.py index 9b0cd6868..245cdc3b6 100644 --- a/test/registered/unit/entrypoints/openai/test_protocol.py +++ b/test/registered/unit/entrypoints/openai/test_protocol.py @@ -124,6 +124,7 @@ class TestChatCompletionRequest(unittest.TestCase): self.assertEqual(request.messages[0].content, "Hello") self.assertEqual(request.temperature, None) # default self.assertFalse(request.stream) # default + self.assertFalse(request.return_sampling_mask) self.assertEqual(request.tool_choice, "none") # default when no tools def test_image_content_hash_validation(self): diff --git a/test/registered/unit/entrypoints/openai/test_serving_chat.py b/test/registered/unit/entrypoints/openai/test_serving_chat.py index 0875c0a30..3140e9faf 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_chat.py +++ b/test/registered/unit/entrypoints/openai/test_serving_chat.py @@ -278,9 +278,12 @@ class ServingChatTestCase(unittest.TestCase): None, ) + self.basic_req.return_sampling_mask = True + self.basic_req.return_meta_info = True adapted, processed = self.chat._convert_to_internal_request(self.basic_req) self.assertIsInstance(adapted, GenerateReqInput) self.assertFalse(adapted.stream) + self.assertTrue(adapted.return_sampling_mask) self.assertEqual(adapted.session_id, "session-1") self.assertEqual(processed, self.basic_req) @@ -295,6 +298,18 @@ class ServingChatTestCase(unittest.TestCase): with self.subTest(field=field), self.assertRaisesRegex(ValueError, field): self.chat._convert_to_internal_request(req, self.fastapi_request) + def test_validate_request_rejects_sampling_mask_without_meta_info(self): + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi?"}], + return_sampling_mask=True, + ) + + self.assertEqual( + self.chat._validate_request(req), + "return_sampling_mask requires return_meta_info=true.", + ) + def test_convert_to_internal_request_rejects_stream_return_meta_info(self): req = ChatCompletionRequest( model="x",