[Sampling] Allow sampling-mask replay with DisallowedTokensLogitsProcessor (#38279)
Co-authored-by: Byron Hsu <24364830+ByronHsu@users.noreply.github.com>
This commit is contained in:
@@ -143,6 +143,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_serving,
|
get_serving,
|
||||||
get_spec,
|
get_spec,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.sampling.custom_logit_processor import supports_sampling_mask
|
||||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
PortArgs,
|
PortArgs,
|
||||||
@@ -1259,6 +1260,17 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
"The server is not configured to enable custom logit processor. "
|
"The server is not configured to enable custom logit processor. "
|
||||||
"Please set `--enable-custom-logit-processor` to enable this feature."
|
"Please set `--enable-custom-logit-processor` to enable this feature."
|
||||||
)
|
)
|
||||||
|
if (
|
||||||
|
obj.return_sampling_mask
|
||||||
|
and obj.custom_logit_processor
|
||||||
|
and not supports_sampling_mask(obj.custom_logit_processor)
|
||||||
|
):
|
||||||
|
# Reject before scheduling so aborted requests cannot execute
|
||||||
|
# unsupported processors during sampling batch preparation.
|
||||||
|
raise ValueError(
|
||||||
|
"return_sampling_mask only supports DisallowedTokensLogitsProcessor "
|
||||||
|
"among custom logit processors."
|
||||||
|
)
|
||||||
|
|
||||||
def _validate_mm_limits(
|
def _validate_mm_limits(
|
||||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
||||||
|
|||||||
@@ -58,6 +58,17 @@ class DisallowedTokensLogitsProcessor(CustomLogitProcessor):
|
|||||||
return logits
|
return logits
|
||||||
|
|
||||||
|
|
||||||
|
def supports_sampling_mask(serialized_processor: str) -> bool:
|
||||||
|
"""Hard exclusion preserves the relative logits needed for mask-based replay."""
|
||||||
|
try:
|
||||||
|
return isinstance(
|
||||||
|
CustomLogitProcessor.from_str(serialized_processor),
|
||||||
|
DisallowedTokensLogitsProcessor,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def _open_thinking_start(ids: list[int], start_id: int, end_id: int) -> int:
|
def _open_thinking_start(ids: list[int], start_id: int, end_id: int) -> int:
|
||||||
"""Return the index of the start token of the currently open thinking block, or -1."""
|
"""Return the index of the start token of the currently open thinking block, or -1."""
|
||||||
for idx in reversed(range(len(ids))):
|
for idx in reversed(range(len(ids))):
|
||||||
|
|||||||
@@ -9,6 +9,11 @@ import torch
|
|||||||
from sglang.srt.layers import sampler as sampler_module
|
from sglang.srt.layers import sampler as sampler_module
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.layers.sampler import Sampler
|
from sglang.srt.layers.sampler import Sampler
|
||||||
|
from sglang.srt.sampling.custom_logit_processor import (
|
||||||
|
DisallowedTokensLogitsProcessor,
|
||||||
|
Qwen3ThinkingBudgetLogitProcessor,
|
||||||
|
)
|
||||||
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
from sglang.srt.utils import is_hip, kill_process_tree
|
from sglang.srt.utils import is_hip, kill_process_tree
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
from sglang.test.test_utils import (
|
from sglang.test.test_utils import (
|
||||||
@@ -30,6 +35,7 @@ _SAMPLING_SEED = 1234
|
|||||||
_SERVER_ARGS = (
|
_SERVER_ARGS = (
|
||||||
"--mem-fraction-static",
|
"--mem-fraction-static",
|
||||||
"0.7",
|
"0.7",
|
||||||
|
"--enable-custom-logit-processor",
|
||||||
)
|
)
|
||||||
_INVALID_SAMPLING_MASK_ERROR = (
|
_INVALID_SAMPLING_MASK_ERROR = (
|
||||||
"top_p-only sampling is valid but can return huge masks in the tail"
|
"top_p-only sampling is valid but can return huge masks in the tail"
|
||||||
@@ -41,6 +47,64 @@ class TestSamplingMaskCapture(CustomTestCase):
|
|||||||
self.sampler = Sampler.__new__(Sampler)
|
self.sampler = Sampler.__new__(Sampler)
|
||||||
torch.nn.Module.__init__(self.sampler)
|
torch.nn.Module.__init__(self.sampler)
|
||||||
|
|
||||||
|
def test_hard_exclusion_replay_in_mixed_batch(self):
|
||||||
|
backends = ["pytorch"] if is_hip() else ["pytorch", "flashinfer"]
|
||||||
|
for backend in backends:
|
||||||
|
with self.subTest(backend=backend):
|
||||||
|
logits = (
|
||||||
|
torch.tensor([[0.3, 0.2, 0.5, 0.15, 0.1]], device="cuda")
|
||||||
|
.log()
|
||||||
|
.repeat(2, 1)
|
||||||
|
)
|
||||||
|
original = logits.clone()
|
||||||
|
info = SamplingBatchInfo(
|
||||||
|
temperatures=torch.ones(2, 1, device="cuda"),
|
||||||
|
top_ps=torch.full((2,), 0.9, device="cuda"),
|
||||||
|
top_ks=torch.full((2,), 3, dtype=torch.int32, device="cuda"),
|
||||||
|
min_ps=torch.zeros(2, device="cuda"),
|
||||||
|
is_all_greedy=False,
|
||||||
|
is_any_greedy=False,
|
||||||
|
need_top_p_sampling=True,
|
||||||
|
need_top_k_sampling=True,
|
||||||
|
need_min_p_sampling=False,
|
||||||
|
vocab_size=5,
|
||||||
|
has_custom_logit_processor=True,
|
||||||
|
custom_params=[{"token_ids": [2]}, None],
|
||||||
|
custom_logit_processor={
|
||||||
|
0: (
|
||||||
|
DisallowedTokensLogitsProcessor(),
|
||||||
|
torch.tensor([True, False], device="cuda"),
|
||||||
|
)
|
||||||
|
},
|
||||||
|
return_sampling_masks=[True, True],
|
||||||
|
)
|
||||||
|
logits = self.sampler._preprocess_logits(logits, info)
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.layers.sampler.get_exec",
|
||||||
|
return_value=SimpleNamespace(
|
||||||
|
kernel=SimpleNamespace(sampling_backend=backend)
|
||||||
|
),
|
||||||
|
):
|
||||||
|
sampled, capture = self.sampler._sample_from_probs(
|
||||||
|
logits.softmax(-1),
|
||||||
|
info,
|
||||||
|
positions=torch.zeros(2, dtype=torch.int64, device="cuda"),
|
||||||
|
simple_sampling_case=False,
|
||||||
|
return_sampling_mask=True,
|
||||||
|
)
|
||||||
|
output = LogitsProcessorOutput(next_token_logits=None)
|
||||||
|
self.sampler._attach_sampling_mask_to_output(
|
||||||
|
output, info, sampled, capture
|
||||||
|
)
|
||||||
|
support = output.next_token_sampling_mask_idx[0]
|
||||||
|
self.assertEqual(set(support), {0, 1, 3})
|
||||||
|
self.assertIn(int(sampled[0]), support)
|
||||||
|
expected = original[0, sampled[0]] - original[0, support].logsumexp(0)
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
output.next_token_sampling_logprobs[0], expected.item(), places=5
|
||||||
|
)
|
||||||
|
self.assertIn(2, output.next_token_sampling_mask_idx[1])
|
||||||
|
|
||||||
@unittest.skipIf(is_hip(), "FlashInfer is not available on ROCm")
|
@unittest.skipIf(is_hip(), "FlashInfer is not available on ROCm")
|
||||||
def test_flashinfer_joint_cutoff_ties_match_capture(self):
|
def test_flashinfer_joint_cutoff_ties_match_capture(self):
|
||||||
batch_size = 256
|
batch_size = 256
|
||||||
@@ -232,12 +296,15 @@ class SamplingMaskTestMixin:
|
|||||||
return_sampling_mask=True,
|
return_sampling_mask=True,
|
||||||
return_logprob=False,
|
return_logprob=False,
|
||||||
top_logprobs_num=0,
|
top_logprobs_num=0,
|
||||||
|
custom_logit_processor=None,
|
||||||
):
|
):
|
||||||
payload = {
|
payload = {
|
||||||
"text": "The capital of France is",
|
"text": "The capital of France is",
|
||||||
"sampling_params": sampling_params,
|
"sampling_params": sampling_params,
|
||||||
"return_sampling_mask": return_sampling_mask,
|
"return_sampling_mask": return_sampling_mask,
|
||||||
}
|
}
|
||||||
|
if custom_logit_processor is not None:
|
||||||
|
payload["custom_logit_processor"] = custom_logit_processor
|
||||||
if return_logprob:
|
if return_logprob:
|
||||||
payload["return_logprob"] = True
|
payload["return_logprob"] = True
|
||||||
payload["top_logprobs_num"] = top_logprobs_num
|
payload["top_logprobs_num"] = top_logprobs_num
|
||||||
@@ -276,6 +343,55 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase):
|
|||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls._launch_server()
|
cls._launch_server()
|
||||||
|
|
||||||
|
def test_disallowed_tokens_with_replay(self):
|
||||||
|
params = {
|
||||||
|
"temperature": 1.0,
|
||||||
|
"top_k": _TOP_K,
|
||||||
|
"top_p": _TOP_P,
|
||||||
|
"max_new_tokens": 1,
|
||||||
|
"ignore_eos": True,
|
||||||
|
}
|
||||||
|
baseline = self._post_generate(params)
|
||||||
|
self.assertEqual(baseline.status_code, 200, baseline.text)
|
||||||
|
# Exclude tokens that actually belong to the unmodified sampling support.
|
||||||
|
blocked = baseline.json()["meta_info"]["output_token_sampling_mask"][0][:2]
|
||||||
|
self.assertTrue(blocked)
|
||||||
|
response = self._post_generate(
|
||||||
|
{**params, "custom_params": {"token_ids": blocked}},
|
||||||
|
return_logprob=True,
|
||||||
|
top_logprobs_num=_TOP_LOGPROBS_NUM,
|
||||||
|
custom_logit_processor=DisallowedTokensLogitsProcessor.to_str(),
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200, response.text)
|
||||||
|
output = response.json()
|
||||||
|
meta = output["meta_info"]
|
||||||
|
token = output["output_ids"][0]
|
||||||
|
mask = meta["output_token_sampling_mask"][0]
|
||||||
|
self.assertTrue(set(mask).isdisjoint(blocked))
|
||||||
|
self.assertIn(token, mask)
|
||||||
|
probs = {
|
||||||
|
int(tid): math.exp(lp) for lp, tid, _ in meta["output_top_logprobs"][0]
|
||||||
|
}
|
||||||
|
expected = math.log(probs[token] / sum(probs[tid] for tid in mask))
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
meta["output_token_sampling_logprobs"][0], expected, delta=1e-2
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_rejected_processors_do_not_break_generation(self):
|
||||||
|
params = {"top_k": _TOP_K, "max_new_tokens": 1}
|
||||||
|
for processor in (
|
||||||
|
Qwen3ThinkingBudgetLogitProcessor.to_str(),
|
||||||
|
"invalid processor",
|
||||||
|
):
|
||||||
|
with self.subTest(processor=processor):
|
||||||
|
response = self._post_generate(params, custom_logit_processor=processor)
|
||||||
|
self.assertEqual(response.status_code, 400, response.text)
|
||||||
|
self.assertIn(
|
||||||
|
"only supports DisallowedTokensLogitsProcessor", response.text
|
||||||
|
)
|
||||||
|
recovery = self._post_generate(params)
|
||||||
|
self.assertEqual(recovery.status_code, 200, recovery.text)
|
||||||
|
|
||||||
def test_generate_returns_sampling_mask(self):
|
def test_generate_returns_sampling_mask(self):
|
||||||
top_p_sampling_masks = self._generate_sampling_masks(
|
top_p_sampling_masks = self._generate_sampling_masks(
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -0,0 +1,83 @@
|
|||||||
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||||
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
|
from sglang.srt.sampling.custom_logit_processor import (
|
||||||
|
CustomLogitProcessor,
|
||||||
|
DisallowedTokensLogitsProcessor,
|
||||||
|
Qwen3ThinkingBudgetLogitProcessor,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestSamplingMaskValidation(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
override = get_context().override_server_args(
|
||||||
|
enable_custom_logit_processor=True
|
||||||
|
)
|
||||||
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
|
self.manager = TokenizerManager.__new__(TokenizerManager)
|
||||||
|
self.manager.context_len = 128
|
||||||
|
self.manager.num_reserved_tokens = 0
|
||||||
|
self.manager.allow_auto_truncate = False
|
||||||
|
self.manager.validate_total_tokens = False
|
||||||
|
self.manager.is_generation = True
|
||||||
|
|
||||||
|
def _validate(self, processor, return_sampling_mask=True):
|
||||||
|
req = GenerateReqInput(
|
||||||
|
input_ids=[1, 2, 3],
|
||||||
|
sampling_params={"top_k": 10},
|
||||||
|
custom_logit_processor=processor,
|
||||||
|
return_sampling_mask=return_sampling_mask,
|
||||||
|
)
|
||||||
|
self.manager._validate_one_request(req, req.input_ids)
|
||||||
|
|
||||||
|
def test_accepts_hard_exclusion(self):
|
||||||
|
class BoundMask(DisallowedTokensLogitsProcessor):
|
||||||
|
def __call__(self, logits, custom_param_list=None):
|
||||||
|
logits[..., [2]] = -float("inf")
|
||||||
|
return logits
|
||||||
|
|
||||||
|
for processor in (
|
||||||
|
None,
|
||||||
|
DisallowedTokensLogitsProcessor.to_str(),
|
||||||
|
BoundMask.to_str(),
|
||||||
|
):
|
||||||
|
with self.subTest(processor=processor):
|
||||||
|
self._validate(processor)
|
||||||
|
|
||||||
|
def test_rejects_unsupported_or_malformed_processors(self):
|
||||||
|
for processor in (
|
||||||
|
Qwen3ThinkingBudgetLogitProcessor.to_str(),
|
||||||
|
"invalid processor",
|
||||||
|
'{"callable": "00"}',
|
||||||
|
):
|
||||||
|
with self.subTest(processor=processor):
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
ValueError, "only supports DisallowedTokensLogitsProcessor"
|
||||||
|
):
|
||||||
|
self._validate(processor)
|
||||||
|
|
||||||
|
def test_disabled_processors_are_not_deserialized(self):
|
||||||
|
with (
|
||||||
|
get_context().override_server_args(enable_custom_logit_processor=False),
|
||||||
|
patch.object(CustomLogitProcessor, "from_str") as deserialize,
|
||||||
|
):
|
||||||
|
with self.assertRaisesRegex(ValueError, "--enable-custom-logit-processor"):
|
||||||
|
self._validate(DisallowedTokensLogitsProcessor.to_str())
|
||||||
|
deserialize.assert_not_called()
|
||||||
|
|
||||||
|
def test_other_processors_still_work_without_sampling_masks(self):
|
||||||
|
self._validate(
|
||||||
|
Qwen3ThinkingBudgetLogitProcessor.to_str(), return_sampling_mask=False
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user