[RL] Expose top-p-only sampling masks (#33593)
This commit is contained in:
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user