[RL] Expose top-p-only sampling masks (#33593)

This commit is contained in:
Nan Jiang
2026-08-14 16:40:27 -07:00
committed by GitHub
parent 90b3db6dd8
commit be804c1b83
6 changed files with 68 additions and 16 deletions
+41 -12
View File
@@ -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):
@@ -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):
@@ -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",