From 5b589ed2e737451846e3d69a91d38e6bff5d02ce Mon Sep 17 00:00:00 2001
From: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
Date: Thu, 7 May 2026 22:21:51 +0100
Subject: [PATCH] feat(constrained): two-phase reasoning grammar +
--enable-strict-thinking (#23953)
---
.../srt/constrained/base_grammar_backend.py | 46 +-
.../sglang/srt/constrained/grammar_manager.py | 40 +-
.../constrained/reasoner_grammar_backend.py | 290 +++++--
.../torch_ops/token_filter_torch_ops.py | 63 ++
.../triton_ops/token_filter_ops.py | 175 ++++
.../srt/constrained/xgrammar_backend.py | 62 +-
python/sglang/srt/environ.py | 3 +
python/sglang/srt/parser/reasoning_parser.py | 32 +-
python/sglang/srt/server_args.py | 9 +
.../constrained/test_base_grammar_backend.py | 20 +-
.../test_e2e_constrained_reasoning.py | 314 +++++++
.../unit/constrained/test_grammar_manager.py | 169 +++-
.../test_reasoner_grammar_backend.py | 774 +++++++++---------
.../unit/constrained/test_token_filter_ops.py | 146 ++++
14 files changed, 1713 insertions(+), 430 deletions(-)
create mode 100644 python/sglang/srt/constrained/torch_ops/token_filter_torch_ops.py
create mode 100644 python/sglang/srt/constrained/triton_ops/token_filter_ops.py
create mode 100644 test/registered/unit/constrained/test_e2e_constrained_reasoning.py
create mode 100644 test/registered/unit/constrained/test_token_filter_ops.py
diff --git a/python/sglang/srt/constrained/base_grammar_backend.py b/python/sglang/srt/constrained/base_grammar_backend.py
index 3704e6d0d..4f907b983 100644
--- a/python/sglang/srt/constrained/base_grammar_backend.py
+++ b/python/sglang/srt/constrained/base_grammar_backend.py
@@ -21,6 +21,7 @@ from typing import Dict, List, Optional, Tuple
import torch
+from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -128,6 +129,8 @@ class InvalidGrammarObject(BaseGrammarObject):
class BaseGrammarBackend:
+ _enable_strict_thinking: bool = False
+
def __init__(self):
self.executor = ThreadPoolExecutor()
self.cache: Dict[Tuple[str, str], BaseGrammarObject] = {}
@@ -136,6 +139,24 @@ class BaseGrammarBackend:
logger.warning(f"Skip unsupported {key_type=}, {key_string=}")
return InvalidGrammarObject()
+ @property
+ def enable_strict_thinking(self):
+ return self._enable_strict_thinking
+
+ @property
+ def is_support_token_filter(self):
+ return False
+
+ def set_token_filter(
+ self, vocab_mask, token_ids, batch_idx, is_allowed=True, reset_vocab_mask=True
+ ):
+ """Set or clear specific tokens in the vocab mask. No-op by default."""
+ pass
+
+ def init_strict_reasoning_grammar(self, reasoning: bool):
+ """Create a grammar object for strict token filtering only. Returns None by default."""
+ return None
+
def dispatch_fallback(self, key_type: str, key_string: str) -> BaseGrammarObject:
"""
This function should not be reached in any case.
@@ -239,6 +260,13 @@ def create_grammar_backend(
any_whitespace=not server_args.constrained_json_disable_any_whitespace,
)
except TokenizerNotSupportedError as e:
+ if server_args.enable_strict_thinking:
+ raise ValueError(
+ f"--enable-strict-thinking requires a grammar backend with "
+ f"token filtering support, but XGrammar failed to initialize: "
+ f"{e}. Cannot fall back to grammar_backend='none' with strict "
+ f"thinking enabled."
+ ) from e
logger.warning(
f"Grammar backend disabled because tokenizer is not supported by XGrammar: {e}. "
"Falling back to grammar_backend='none'. "
@@ -255,6 +283,13 @@ def create_grammar_backend(
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
)
elif name == "none":
+ if server_args.enable_strict_thinking:
+ raise ValueError(
+ "--enable-strict-thinking requires a grammar backend that supports "
+ "token filtering, but grammar_backend='none' was specified. Use "
+ "--grammar-backend xgrammar or another backend that supports token "
+ "filtering."
+ )
return None
else:
raise ValueError(f"Invalid grammar backend: {name}")
@@ -264,6 +299,15 @@ def create_grammar_backend(
ReasonerGrammarBackend,
)
- grammar_backend = ReasonerGrammarBackend(grammar_backend, think_end_id)
+ reasoning_parser = ReasoningParser(
+ model_type=server_args.reasoning_parser, stream_reasoning=False
+ )
+
+ grammar_backend = ReasonerGrammarBackend(
+ grammar_backend,
+ reasoning_parser,
+ tokenizer,
+ enable_strict_thinking=server_args.enable_strict_thinking,
+ )
return grammar_backend
diff --git a/python/sglang/srt/constrained/grammar_manager.py b/python/sglang/srt/constrained/grammar_manager.py
index 8b1e79658..5442cb5d3 100644
--- a/python/sglang/srt/constrained/grammar_manager.py
+++ b/python/sglang/srt/constrained/grammar_manager.py
@@ -11,6 +11,7 @@ from sglang.srt.constrained.base_grammar_backend import (
InvalidGrammarObject,
create_grammar_backend,
)
+from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
from sglang.srt.environ import envs
if TYPE_CHECKING:
@@ -37,6 +38,12 @@ class GrammarManager:
else:
self.grammar_backend = None
+ self._enable_strict_thinking = (
+ self.grammar_backend.enable_strict_thinking
+ if self.grammar_backend is not None
+ else False
+ )
+
self.grammar_sync_group = scheduler.dp_tp_cpu_group
self.grammar_sync_size = scheduler.dp_tp_group.world_size
self.grammar_sync_entry = scheduler.dp_tp_group.first_rank
@@ -65,6 +72,20 @@ class GrammarManager:
req.grammar.cancel()
req.set_finish_with_abort("Aborted by AbortReq.")
+ def _get_request_thinking_budget(self, req: Req) -> int | None:
+ custom_params = req.sampling_params.custom_params
+ if not isinstance(custom_params, dict):
+ return None
+ thinking_budget = custom_params.get("thinking_budget")
+ return thinking_budget if isinstance(thinking_budget, int) else None
+
+ def _apply_request_reasoning_budget(self, req: Req) -> None:
+ thinking_budget = self._get_request_thinking_budget(req)
+ if thinking_budget is None:
+ return
+ if isinstance(req.grammar, ReasonerGrammarObject):
+ req.grammar.max_think_tokens = thinking_budget
+
def process_req_with_grammar(self, req: Req) -> bool:
# Init grammar cache for this request
add_to_grammar_queue = False
@@ -103,6 +124,15 @@ class GrammarManager:
f"Failed to compile {key[0]} grammar: {value.error_message}"
)
req.set_finish_with_abort(error_msg)
+ else:
+ self._apply_request_reasoning_budget(req)
+ elif self._enable_strict_thinking:
+ grammar_obj = self.grammar_backend.init_strict_reasoning_grammar(
+ req.require_reasoning
+ )
+ if grammar_obj is not None:
+ req.grammar = grammar_obj
+ self._apply_request_reasoning_budget(req)
if add_to_grammar_queue:
self.grammar_queue.append(req)
@@ -177,8 +207,16 @@ class GrammarManager:
continue
assert isinstance(req.grammar, futures.Future) and req.grammar_key
- req.grammar = req.grammar.result()
+ try:
+ req.grammar = req.grammar.result()
+ except Exception as e:
+ logger.error(
+ f"Grammar compilation raised an exception: {e}, "
+ f"grammar_key={req.grammar_key}"
+ )
+ req.grammar = InvalidGrammarObject(f"Grammar compilation failed: {e}")
self.grammar_backend.set_cache(req.grammar_key, req.grammar.copy())
+ self._apply_request_reasoning_budget(req)
if isinstance(req.grammar, InvalidGrammarObject):
error_msg = f"Failed to compile {req.grammar_key[0]} grammar: {req.grammar.error_message}"
req.set_finish_with_abort(error_msg)
diff --git a/python/sglang/srt/constrained/reasoner_grammar_backend.py b/python/sglang/srt/constrained/reasoner_grammar_backend.py
index d204bdd9e..57d0f65d5 100644
--- a/python/sglang/srt/constrained/reasoner_grammar_backend.py
+++ b/python/sglang/srt/constrained/reasoner_grammar_backend.py
@@ -13,9 +13,14 @@
# ==============================================================================
"""The baseclass of a backend for reasoner grammar-guided constrained decoding."""
-from typing import List, Optional, Tuple
+import logging
+from typing import List, Optional, Tuple, Union
import torch
+from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast
+
+from sglang.srt.environ import envs
+from sglang.srt.parser.reasoning_parser import ReasoningParser
from .base_grammar_backend import (
BaseGrammarBackend,
@@ -23,102 +28,291 @@ from .base_grammar_backend import (
InvalidGrammarObject,
)
+logger = logging.getLogger(__name__)
+
class ReasonerGrammarObject(BaseGrammarObject):
- def __init__(self, grammar: BaseGrammarObject, think_end_id: int):
+ """Wraps a grammar object to handle reasoning (think/generation) phases.
+
+ State machine (must call maybe_init_reasoning before use):
+ THINKING (tokens_in_think >= 0, tokens_after_end == -1)
+ -> grammar not consulted, optional token filtering
+ GENERATION (tokens_after_end >= 0)
+ -> grammar consulted for accept/fill/rollback
+
+ When enable_token_filter=True (strict mode), fill_vocab_mask filters
+ excluded tokens during THINKING and enforces max_think_tokens budget.
+ When the budget is exhausted, only think_end_id is allowed, forcing the
+ model to exit the thinking phase.
+ When enable_token_filter=False (non-strict mode), fill_vocab_mask is
+ a no-op during THINKING.
+ """
+
+ def __init__(
+ self,
+ grammar: Optional[BaseGrammarObject],
+ think_end_id: int,
+ think_excluded_token_ids: Optional[List[int]] = None,
+ max_think_tokens: int = -1,
+ enable_token_filter: bool = False,
+ token_filter_fn=None,
+ allocate_vocab_mask_fn=None,
+ move_vocab_mask_fn=None,
+ apply_vocab_mask_fn=None,
+ ):
super().__init__()
self.grammar = grammar
self.think_end_id = think_end_id
- # -1 means thinking has not ended yet
- # 0 means just ended thinking in the last token
- # + means number of tokens after thinking ended
- self.tokens_after_think_end = -1
+ self.think_excluded_token_ids = think_excluded_token_ids
+ self.max_think_tokens = max_think_tokens
+ self.enable_token_filter = enable_token_filter
+ self.token_filter_fn = token_filter_fn
+ self.allocate_vocab_mask_fn = allocate_vocab_mask_fn
+ self.move_vocab_mask_fn = move_vocab_mask_fn
+ self.apply_vocab_mask_fn = apply_vocab_mask_fn
+ self._think_end_id_list = [think_end_id]
+
+ self.tokens_in_think = -1
+ self.tokens_after_end = -1
def maybe_init_reasoning(self, reasoning: bool):
- self.tokens_after_think_end = -1 if reasoning else 0
+ if reasoning:
+ self.tokens_in_think = 0
+ else:
+ self.tokens_in_think = -1
+ self.tokens_after_end = 0
- def transfer_state(self, token: int) -> int:
- if self.tokens_after_think_end == -1 and token == self.think_end_id:
- self.tokens_after_think_end = 0
- elif self.tokens_after_think_end >= 0:
- self.tokens_after_think_end += 1
+ def _is_thinking(self):
+ return self.tokens_in_think >= 0 and self.tokens_after_end == -1
+
+ def _is_generation(self):
+ return self.tokens_after_end >= 0
+
+ def transfer_state(self, token: int) -> None:
+ if self._is_thinking():
+ if token == self.think_end_id:
+ self.tokens_after_end = 0
+ else:
+ self.tokens_in_think += 1
+ elif self._is_generation():
+ self.tokens_after_end += 1
def rollback_state(self):
- if self.tokens_after_think_end == 0:
- self.tokens_after_think_end = -1
- elif self.tokens_after_think_end > 0:
- self.tokens_after_think_end -= 1
+ if self._is_thinking():
+ if self.tokens_in_think > 0:
+ self.tokens_in_think -= 1
+ elif self._is_generation():
+ if self.tokens_after_end == 0:
+ self.tokens_after_end = -1
+ elif self.tokens_after_end > 0:
+ self.tokens_after_end -= 1
def accept_token(self, token: int):
- if self.tokens_after_think_end >= 0:
+ if self._is_generation() and self.grammar is not None:
self.grammar.accept_token(token)
self.transfer_state(token)
def is_terminated(self):
- return self.grammar.is_terminated()
+ if self.grammar is not None:
+ return self.grammar.is_terminated()
+ return False
def rollback(self, k):
- steps_after_think = min(k, self.tokens_after_think_end)
- if steps_after_think > 0:
- self.grammar.rollback(steps_after_think)
-
+ if self.grammar is not None:
+ steps_after = min(k, max(0, self.tokens_after_end))
+ if steps_after > 0:
+ self.grammar.rollback(steps_after)
for _ in range(k):
self.rollback_state()
- def allocate_vocab_mask(
- self, vocab_size: int, batch_size: int, device
- ) -> torch.Tensor:
- return self.grammar.allocate_vocab_mask(vocab_size, batch_size, device)
+ def _can_think_more(self):
+ return self.max_think_tokens < 0 or self.tokens_in_think < self.max_think_tokens
+
+ def _do_token_filter(self, vocab_mask, token_ids, idx, is_allowed=True):
+ if self.token_filter_fn is not None:
+ self.token_filter_fn(vocab_mask, token_ids, idx, is_allowed)
def fill_vocab_mask(self, vocab_mask: torch.Tensor, idx: int) -> None:
- if self.tokens_after_think_end >= 0:
+ if self._is_thinking():
+ if not self.enable_token_filter:
+ return
+ if self._can_think_more():
+ self._do_token_filter(
+ vocab_mask, self.think_excluded_token_ids, idx, is_allowed=False
+ )
+ else:
+ self._do_token_filter(
+ vocab_mask, self._think_end_id_list, idx, is_allowed=True
+ )
+ return
+ if self._is_generation() and self.grammar is not None:
self.grammar.fill_vocab_mask(vocab_mask, idx)
- def move_vocab_mask(self, vocab_mask: torch.Tensor, device) -> torch.Tensor:
- return self.grammar.move_vocab_mask(vocab_mask, device)
+ def allocate_vocab_mask(self, vocab_size, batch_size, device):
+ if self.grammar is not None:
+ return self.grammar.allocate_vocab_mask(vocab_size, batch_size, device)
+ if self.allocate_vocab_mask_fn is not None:
+ return self.allocate_vocab_mask_fn(vocab_size, batch_size, device)
+ return None
+
+ def move_vocab_mask(self, vocab_mask, device):
+ if self.grammar is not None:
+ return self.grammar.move_vocab_mask(vocab_mask, device)
+ if self.move_vocab_mask_fn is not None:
+ return self.move_vocab_mask_fn(vocab_mask, device)
+ return vocab_mask
@property
def apply_vocab_mask(self):
- return self.grammar.apply_vocab_mask
+ if self.grammar is not None:
+ return self.grammar.apply_vocab_mask
+ return self.apply_vocab_mask_fn
- def copy(self) -> BaseGrammarObject:
- return ReasonerGrammarObject(self.grammar.copy(), self.think_end_id)
+ def copy(self):
+ new_obj = ReasonerGrammarObject(
+ self.grammar.copy() if self.grammar is not None else None,
+ self.think_end_id,
+ self.think_excluded_token_ids,
+ self.max_think_tokens,
+ self.enable_token_filter,
+ self.token_filter_fn,
+ self.allocate_vocab_mask_fn,
+ self.move_vocab_mask_fn,
+ self.apply_vocab_mask_fn,
+ )
+ new_obj.tokens_in_think = self.tokens_in_think
+ new_obj.tokens_after_end = self.tokens_after_end
+ new_obj._finished = self._finished
+ return new_obj
@property
def finished(self):
- return self.grammar.finished
+ if self.grammar is not None:
+ return self.grammar.finished
+ return self._finished
@finished.setter
def finished(self, finished):
- self.grammar.finished = finished
+ if self.grammar is not None:
+ self.grammar.finished = finished
+ else:
+ self._finished = finished
def try_jump_forward(self, tokenizer):
- return self.grammar.try_jump_forward(tokenizer)
+ if self.grammar is not None:
+ return self.grammar.try_jump_forward(tokenizer)
+ return None
def jump_forward_str_state(self, helper):
- return self.grammar.jump_forward_str_state(helper)
+ if self.grammar is not None:
+ return self.grammar.jump_forward_str_state(helper)
+ return None
- def jump_and_retokenize(
- self, old_output_ids: List[int], new_output_ids: List[int], next_state: int
- ):
- return self.grammar.jump_and_retokenize(
- old_output_ids, new_output_ids, next_state
- )
+ def jump_and_retokenize(self, old_output_ids, new_output_ids, next_state):
+ if self.grammar is not None:
+ return self.grammar.jump_and_retokenize(
+ old_output_ids, new_output_ids, next_state
+ )
class ReasonerGrammarBackend(BaseGrammarBackend):
- def __init__(self, grammar_backend: BaseGrammarBackend, think_end_id):
+ def __init__(
+ self,
+ grammar_backend: BaseGrammarBackend,
+ reasoning_parser: ReasoningParser,
+ tokenizer: Union[PreTrainedTokenizer, PreTrainedTokenizerFast],
+ enable_strict_thinking: bool = False,
+ ):
super().__init__()
self.grammar_backend = grammar_backend
- self.think_end_id = think_end_id
+ think_end_ids = tokenizer.encode(
+ reasoning_parser.detector.think_end_token, add_special_tokens=False
+ )
+ if not think_end_ids:
+ raise ValueError(
+ f"think_end_token '{reasoning_parser.detector.think_end_token}' "
+ f"could not be encoded by the tokenizer."
+ )
+ if len(think_end_ids) != 1:
+ raise ValueError(
+ f"think_end_token '{reasoning_parser.detector.think_end_token}' "
+ "must encode to exactly one token for constrained reasoning."
+ )
+ self.think_end_id = think_end_ids[0]
+ self._enable_strict_thinking = enable_strict_thinking
+ self.think_excluded_token_ids = self._get_think_excluded_token_ids(
+ reasoning_parser, tokenizer
+ )
+ self.max_think_tokens = envs.SGLANG_MAX_THINK_TOKENS.get()
+ if (
+ self.enable_strict_thinking
+ and self.think_excluded_token_ids is not None
+ and not self.grammar_backend.is_support_token_filter
+ ):
+ raise ValueError(
+ "Strict reasoning format requested but the grammar backend does not "
+ "support token filtering. Use a grammar backend that supports token "
+ "filtering (e.g., xgrammar) or disable strict reasoning mode."
+ )
+ self.enable_token_filter = (
+ self.enable_strict_thinking
+ and self.think_excluded_token_ids is not None
+ and self.grammar_backend.is_support_token_filter
+ )
+ self._token_filter_fn = (
+ self.grammar_backend.set_token_filter if self.enable_token_filter else None
+ )
+
+ def _get_think_excluded_token_ids(
+ self,
+ reasoning_parser: ReasoningParser,
+ tokenizer: Union[PreTrainedTokenizer, PreTrainedTokenizerFast],
+ ) -> Optional[List[int]]:
+ excluded_ids = []
+ if (not self.enable_strict_thinking) or (
+ not reasoning_parser.detector.think_excluded_tokens
+ ):
+ return None
+ for token in reasoning_parser.detector.think_excluded_tokens:
+ new_ids = tokenizer.encode(token, add_special_tokens=False)
+ if not new_ids:
+ raise ValueError(
+ f"think_excluded_token '{token}' could not be encoded by the "
+ f"tokenizer. All excluded tokens must be encodable for strict "
+ f"reasoning mode to function correctly."
+ )
+ excluded_ids += new_ids
+ return excluded_ids
+
+ def _make_grammar_object(
+ self, grammar: Optional[BaseGrammarObject], reasoning: bool
+ ) -> ReasonerGrammarObject:
+ obj = ReasonerGrammarObject(
+ grammar=grammar,
+ think_end_id=self.think_end_id,
+ think_excluded_token_ids=self.think_excluded_token_ids,
+ max_think_tokens=self.max_think_tokens,
+ enable_token_filter=self.enable_token_filter,
+ token_filter_fn=self._token_filter_fn,
+ allocate_vocab_mask_fn=self.grammar_backend.allocate_vocab_mask,
+ move_vocab_mask_fn=self.grammar_backend.move_vocab_mask,
+ apply_vocab_mask_fn=self.grammar_backend.apply_vocab_mask,
+ )
+ obj.maybe_init_reasoning(reasoning)
+ return obj
+
+ def init_strict_reasoning_grammar(
+ self, reasoning: bool
+ ) -> Optional[BaseGrammarObject]:
+ """Create a grammar object for strict token filtering only (no inner grammar)."""
+ if not self.enable_strict_thinking:
+ return None
+ return self._make_grammar_object(None, reasoning)
def _init_value_dispatch(
self, key: Tuple[str, str], reasoning: bool
) -> Optional[BaseGrammarObject]:
ret = self.grammar_backend._init_value_dispatch(key, reasoning)
- # avoid wrapping invalid grammar, so that the scheduler can detect it
if ret is None or isinstance(ret, InvalidGrammarObject):
return ret
- obj = ReasonerGrammarObject(ret, self.think_end_id)
- obj.maybe_init_reasoning(reasoning)
- return obj
+ return self._make_grammar_object(ret, reasoning)
diff --git a/python/sglang/srt/constrained/torch_ops/token_filter_torch_ops.py b/python/sglang/srt/constrained/torch_ops/token_filter_torch_ops.py
new file mode 100644
index 000000000..f12d62b46
--- /dev/null
+++ b/python/sglang/srt/constrained/torch_ops/token_filter_torch_ops.py
@@ -0,0 +1,63 @@
+# Copyright 2026 SGLang Team
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+# ==============================================================================
+"""Torch fallback for token filter operations (non-CUDA devices and HIP).
+
+Sets or clears specific bits in an int32 bitmask by token ID. The token list
+is typically tiny (< 10 entries); aggregation is done in Python with the actual
+bitmask operations using torch tensor indexing.
+"""
+
+import ctypes
+from typing import List
+
+import torch
+
+
+def set_token_filter_torch(
+ vocab_mask: torch.Tensor,
+ token_ids: List[int],
+ batch_idx: int,
+ is_allowed: bool = True,
+ reset_vocab_mask: bool = True,
+):
+ if reset_vocab_mask:
+ vocab_mask[batch_idx].fill_(-1 if (not is_allowed) else 0)
+
+ if not token_ids:
+ return
+
+ # Aggregate bit masks per int32 element to handle duplicate indices.
+ aggregated: dict[int, int] = {}
+ for token_id in token_ids:
+ element_idx = token_id // 32
+ bit_idx = token_id % 32
+ aggregated[element_idx] = aggregated.get(element_idx, 0) | (1 << bit_idx)
+
+ row = vocab_mask[batch_idx]
+ element_indices = torch.tensor(
+ list(aggregated.keys()), dtype=torch.long, device=row.device
+ )
+ bitmasks = torch.tensor(
+ [
+ ctypes.c_int32(mask if is_allowed else ~mask).value
+ for mask in aggregated.values()
+ ],
+ dtype=row.dtype,
+ device=row.device,
+ )
+
+ if is_allowed:
+ row[element_indices] = torch.bitwise_or(row[element_indices], bitmasks)
+ else:
+ row[element_indices] = torch.bitwise_and(row[element_indices], bitmasks)
diff --git a/python/sglang/srt/constrained/triton_ops/token_filter_ops.py b/python/sglang/srt/constrained/triton_ops/token_filter_ops.py
new file mode 100644
index 000000000..f8d6b6cfc
--- /dev/null
+++ b/python/sglang/srt/constrained/triton_ops/token_filter_ops.py
@@ -0,0 +1,175 @@
+# Copyright 2026 SGLang Team
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+# ==============================================================================
+"""Triton kernels for token filter operations."""
+
+from collections import OrderedDict
+from typing import List
+
+import torch
+import triton
+import triton.language as tl
+
+from sglang.srt.utils import get_device_core_count
+
+
+@triton.jit
+def reset_vocab_mask_kernel(
+ vocab_mask_ptr,
+ batch_idx: int,
+ num_elements: int,
+ reset_value: tl.constexpr,
+):
+ """Reset the vocab mask for a specific batch index to a given value.
+
+ Parameters
+ ----------
+ vocab_mask_ptr : tl.tensor
+ Pointer to the vocab mask tensor.
+
+ batch_idx : int
+ The batch index to reset.
+
+ num_elements : int
+ Number of int32 elements in the vocab mask for each batch.
+
+ reset_value : int
+ The value to reset the vocab mask to (typically -1 or 0).
+ """
+ pid = tl.program_id(0)
+ num_threads = tl.num_programs(0)
+
+ for i in tl.range(pid, num_elements, num_threads):
+ offset = batch_idx * num_elements + i
+ tl.store(vocab_mask_ptr + offset, reset_value)
+
+
+@triton.jit
+def set_token_filter_batch_kernel(
+ vocab_mask_ptr,
+ token_ids_ptr,
+ batch_idx: int,
+ num_tokens: int,
+ num_elements: int,
+ is_allowed: tl.constexpr,
+):
+ """Set or clear specific tokens in the vocab mask for a batch.
+
+ Each token ID maps to a specific bit in the int32 bitmask array.
+ The kernel sets or clears those bits using atomic operations.
+
+ Parameters
+ ----------
+ vocab_mask_ptr : tl.tensor
+ Pointer to the vocab mask tensor.
+
+ token_ids_ptr : tl.tensor
+ Pointer to the token IDs to set/clear.
+
+ batch_idx : int
+ The batch index to modify.
+
+ num_tokens : int
+ Number of tokens to process.
+
+ num_elements : int
+ Number of int32 elements in the vocab mask for each batch.
+
+ is_allowed : bool
+ If True, set the bit to 1 (allow token).
+ If False, clear the bit to 0 (block token).
+ """
+ pid = tl.program_id(0)
+ num_threads = tl.num_programs(0)
+
+ for i in tl.range(pid, num_tokens, num_threads):
+ token_id = tl.load(token_ids_ptr + i)
+ element_idx = token_id // 32
+ bit_idx = token_id % 32
+
+ offset = batch_idx * num_elements + element_idx
+
+ if is_allowed:
+ tl.atomic_or(vocab_mask_ptr + offset, 1 << bit_idx)
+ else:
+ tl.atomic_and(vocab_mask_ptr + offset, ~(1 << bit_idx))
+
+
+_cached_num_sms = None
+_cached_token_id_tensors: OrderedDict[tuple[int, tuple[int, ...]], torch.Tensor] = (
+ OrderedDict()
+)
+_MAX_TOKEN_ID_TENSOR_CACHE_SIZE = 32
+
+
+def _compute_grid(work_items: int):
+ global _cached_num_sms
+ if _cached_num_sms is None:
+ _cached_num_sms = get_device_core_count()
+ if _cached_num_sms > 0:
+ return (min(_cached_num_sms, work_items),)
+ return (work_items,)
+
+
+def _get_cached_token_ids_tensor(
+ token_ids: List[int], device: torch.device
+) -> torch.Tensor:
+ key = (device.index or 0, tuple(token_ids))
+ cached = _cached_token_id_tensors.get(key)
+ if cached is not None:
+ _cached_token_id_tensors.move_to_end(key)
+ return cached
+
+ token_ids_tensor = torch.tensor(token_ids, dtype=torch.int32, device=device)
+ _cached_token_id_tensors[key] = token_ids_tensor
+ if len(_cached_token_id_tensors) > _MAX_TOKEN_ID_TENSOR_CACHE_SIZE:
+ _cached_token_id_tensors.popitem(last=False)
+ return token_ids_tensor
+
+
+def set_token_filter_triton(
+ vocab_mask: torch.Tensor,
+ token_ids: List[int],
+ batch_idx: int,
+ is_allowed: bool = True,
+ reset_vocab_mask: bool = True,
+):
+ """Set or clear specific tokens in the vocab mask using Triton."""
+ assert vocab_mask.device.type == "cuda"
+
+ num_elements = vocab_mask.shape[1]
+
+ if reset_vocab_mask:
+ reset_value = 0 if is_allowed else -1
+ reset_vocab_mask_kernel[_compute_grid(num_elements)](
+ vocab_mask,
+ batch_idx,
+ num_elements,
+ reset_value,
+ num_warps=4,
+ )
+
+ if not token_ids:
+ return
+
+ num_tokens = len(token_ids)
+ token_ids_tensor = _get_cached_token_ids_tensor(token_ids, vocab_mask.device)
+ set_token_filter_batch_kernel[_compute_grid(num_tokens)](
+ vocab_mask,
+ token_ids_tensor,
+ batch_idx,
+ num_tokens,
+ num_elements,
+ is_allowed,
+ num_warps=4,
+ )
diff --git a/python/sglang/srt/constrained/xgrammar_backend.py b/python/sglang/srt/constrained/xgrammar_backend.py
index 542e920ce..46ce3c305 100644
--- a/python/sglang/srt/constrained/xgrammar_backend.py
+++ b/python/sglang/srt/constrained/xgrammar_backend.py
@@ -49,6 +49,10 @@ else:
apply_token_bitmask_inplace_triton,
)
+from sglang.srt.constrained.torch_ops.token_filter_torch_ops import (
+ set_token_filter_torch,
+)
+from sglang.srt.constrained.triton_ops.token_filter_ops import set_token_filter_triton
logger = logging.getLogger(__name__)
MAX_ROLLBACK_TOKENS = 200
@@ -62,7 +66,7 @@ class XGrammarGrammar(BaseGrammarObject):
vocab_size: int,
ctx: CompiledGrammar,
override_stop_tokens: Optional[Union[List[int], int]],
- key_string: Optional[str] = None, # TODO (sk): for debugging, remove later
+ key_string: Optional[str] = None,
grammar_stats: Optional[GrammarStats] = GrammarStats(),
) -> None:
super().__init__()
@@ -162,7 +166,14 @@ class XGrammarGrammar(BaseGrammarObject):
self.matcher.rollback(len(old_output_ids) - k)
for i in range(k, len(new_output_ids)):
- assert self.matcher.accept_token(new_output_ids[i])
+ if not self.matcher.accept_token(new_output_ids[i]):
+ raise ValueError(
+ f"Token not accepted during retokenization: {new_output_ids[i]} "
+ f"at position {i}\n"
+ f"Old output IDs: {old_output_ids}\n"
+ f"New output IDs: {new_output_ids}\n"
+ f"Key string: {self.key_string}"
+ )
def __repr__(self):
return f"XGrammarGrammar({self.key_string=}, {self.accepted_tokens=}, {self.current_token=})"
@@ -211,6 +222,53 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
self.override_stop_tokens = override_stop_tokens
self.any_whitespace = any_whitespace
+ @property
+ def is_support_token_filter(self):
+ return True
+
+ @staticmethod
+ def allocate_vocab_mask(vocab_size: int, batch_size: int, device) -> torch.Tensor:
+ return allocate_token_bitmask(batch_size, vocab_size)
+
+ @staticmethod
+ def move_vocab_mask(vocab_mask: torch.Tensor, device) -> torch.Tensor:
+ return vocab_mask.to(device, non_blocking=True)
+
+ @staticmethod
+ def apply_vocab_mask(logits: torch.Tensor, vocab_mask: torch.Tensor) -> None:
+ if logits.device.type in {"cuda", "npu", "xpu", "musa"}:
+ if _is_hip:
+ apply_token_bitmask_inplace_cuda(logits, vocab_mask)
+ else:
+ apply_token_bitmask_inplace_triton(logits, vocab_mask)
+ else:
+ raise RuntimeError(f"Unsupported device: {logits.device.type}")
+
+ @staticmethod
+ def set_token_filter(
+ vocab_mask: torch.Tensor,
+ token_ids: List[int],
+ batch_idx: int,
+ is_allowed: bool = True,
+ reset_vocab_mask: bool = True,
+ ):
+ if _is_hip or (vocab_mask.device.type != "cuda"):
+ set_token_filter_torch(
+ vocab_mask,
+ token_ids,
+ batch_idx,
+ is_allowed=is_allowed,
+ reset_vocab_mask=reset_vocab_mask,
+ )
+ else:
+ set_token_filter_triton(
+ vocab_mask,
+ token_ids,
+ batch_idx,
+ is_allowed=is_allowed,
+ reset_vocab_mask=reset_vocab_mask,
+ )
+
@staticmethod
def _sanitize_structural_format(structural_format):
"""Recursively replace missing json_schema fields with an empty schema."""
diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py
index 55a0245eb..53bc59279 100644
--- a/python/sglang/srt/environ.py
+++ b/python/sglang/srt/environ.py
@@ -507,6 +507,9 @@ class Envs:
# Tool-Call behavior
SGLANG_TOOL_STRICT_LEVEL = EnvInt(ToolStrictLevel.OFF)
+ # Think tokens budget: negative means unlimited, >= 0 caps thinking tokens
+ SGLANG_MAX_THINK_TOKENS = EnvInt(-1)
+
# Ngram
SGLANG_NGRAM_FORCE_GREEDY_VERIFY = EnvBool(False)
diff --git a/python/sglang/srt/parser/reasoning_parser.py b/python/sglang/srt/parser/reasoning_parser.py
index 6c80033c2..9ad15c8d1 100644
--- a/python/sglang/srt/parser/reasoning_parser.py
+++ b/python/sglang/srt/parser/reasoning_parser.py
@@ -1,4 +1,4 @@
-from typing import Dict, Optional, Tuple, Type
+from typing import Dict, List, Optional, Tuple, Type
from sglang.srt.entrypoints.openai.protocol import ChatCompletionRequest
from sglang.srt.parser.harmony_parser import HarmonyParser
@@ -23,6 +23,7 @@ class BaseReasoningFormatDetector:
self,
think_start_token: str,
think_end_token: str,
+ think_excluded_tokens: Optional[List[str]] = None,
force_reasoning: bool = False,
stream_reasoning: bool = True,
tool_start_token: Optional[str] = None,
@@ -33,6 +34,7 @@ class BaseReasoningFormatDetector:
):
self.think_start_token = think_start_token
self.think_end_token = think_end_token
+ self.think_excluded_tokens = think_excluded_tokens
self.tool_start_token = tool_start_token
self.force_reasoning = force_reasoning
self._in_reasoning = force_reasoning
@@ -242,9 +244,16 @@ class Qwen3Detector(BaseReasoningFormatDetector):
continue_final_message: bool = False,
previous_content: str = "",
):
+ think_excluded_tokens = [
+ "",
+ "",
+ "<|im_end|>",
+ "<|endoftext|>",
+ ]
super().__init__(
"",
"",
+ think_excluded_tokens=think_excluded_tokens,
force_reasoning=force_reasoning,
stream_reasoning=stream_reasoning,
continue_final_message=continue_final_message,
@@ -297,9 +306,22 @@ class KimiK2Detector(BaseReasoningFormatDetector):
continue_final_message: bool = False,
previous_content: str = "",
):
+ think_excluded_tokens = [
+ "",
+ "<|tool_calls_section_begin|>",
+ "<|tool_call_begin|>",
+ "<|tool_call_argument_begin|>",
+ "<|tool_call_section_end|>",
+ "<|tool_call_end|>",
+ "[EOS]",
+ "<|im_end|>",
+ "<|end_header_id|>",
+ "[EOT]",
+ ]
super().__init__(
"",
"",
+ think_excluded_tokens=think_excluded_tokens,
force_reasoning=force_reasoning,
stream_reasoning=stream_reasoning,
tool_start_token="<|tool_calls_section_begin|>",
@@ -323,9 +345,17 @@ class Glm45Detector(BaseReasoningFormatDetector):
"""
def __init__(self, stream_reasoning: bool = True, force_reasoning: bool = False):
+ think_excluded_tokens = [
+ "",
+ "",
+ "",
+ "<|user|>",
+ "<|endoftext|>",
+ ]
super().__init__(
"",
"",
+ think_excluded_tokens=think_excluded_tokens,
force_reasoning=force_reasoning,
stream_reasoning=stream_reasoning,
tool_start_token="",
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index 40ffb001c..de7f06468 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -495,6 +495,7 @@ class ServerArgs:
enable_cache_report: bool = False
reasoning_parser: Optional[str] = None
strip_thinking_cache: bool = False
+ enable_strict_thinking: bool = False
tool_call_parser: Optional[str] = None
tool_server: Optional[str] = None
sampling_defaults: str = "model"
@@ -5127,6 +5128,14 @@ class ServerArgs:
"radix tree on finish; keep only the prompt prefix. Opt-in: changes "
"cache contents.",
)
+ parser.add_argument(
+ "--enable-strict-thinking",
+ action="store_true",
+ default=ServerArgs.enable_strict_thinking,
+ help="Enable strict token filtering during the thinking phase. "
+ "Blocks model-specific excluded tokens (e.g., tool call markers) "
+ "during reasoning. Requires a grammar backend that supports token filtering.",
+ )
tool_call_parser_choices = list(FunctionCallParser.ToolCallParserEnum.keys())
parser.add_argument(
"--tool-call-parser",
diff --git a/test/registered/unit/constrained/test_base_grammar_backend.py b/test/registered/unit/constrained/test_base_grammar_backend.py
index 1f3a7a4f9..191fb4b3a 100644
--- a/test/registered/unit/constrained/test_base_grammar_backend.py
+++ b/test/registered/unit/constrained/test_base_grammar_backend.py
@@ -269,10 +269,13 @@ class TestCreateGrammarBackend(unittest.TestCase):
GRAMMAR_BACKEND_REGISTRY.clear()
GRAMMAR_BACKEND_REGISTRY.update(self._saved)
- def _make_server_args(self, backend="none", reasoning_parser=None):
+ def _make_server_args(
+ self, backend="none", reasoning_parser=None, enable_strict_thinking=False
+ ):
args = MagicMock()
args.grammar_backend = backend
args.reasoning_parser = reasoning_parser
+ args.enable_strict_thinking = enable_strict_thinking
args.constrained_json_whitespace_pattern = None
args.constrained_json_disable_any_whitespace = False
return args
@@ -282,6 +285,11 @@ class TestCreateGrammarBackend(unittest.TestCase):
result = create_grammar_backend(args, None, 32000)
self.assertIsNone(result)
+ def test_none_backend_with_strict_thinking_raises(self):
+ args = self._make_server_args("none", enable_strict_thinking=True)
+ with self.assertRaisesRegex(ValueError, "enable-strict-thinking"):
+ create_grammar_backend(args, None, 32000)
+
def test_invalid_backend_raises(self):
args = self._make_server_args("nonexistent_backend")
with self.assertRaises(ValueError):
@@ -316,7 +324,7 @@ class TestCreateGrammarBackend(unittest.TestCase):
mock_inner = MagicMock(spec=BaseGrammarBackend)
register_grammar_backend("inner_r", lambda *a: mock_inner)
- args = self._make_server_args("inner_r", reasoning_parser="deepseek")
+ args = self._make_server_args("inner_r", reasoning_parser="deepseek-r1")
tokenizer = MagicMock()
result = create_grammar_backend(args, tokenizer, 32000)
@@ -382,13 +390,15 @@ class TestCreateGrammarBackend(unittest.TestCase):
)
mock_backend = MagicMock(spec=BaseGrammarBackend)
+ mock_backend.is_support_token_filter = False
mock_outlines_cls.return_value = mock_backend
- args = self._make_server_args("outlines", reasoning_parser="deepseek")
+ args = self._make_server_args("outlines", reasoning_parser="deepseek-r1")
tokenizer = MagicMock()
+ # encode must return a single-token list for think_start/end tokens
+ tokenizer.encode.return_value = [42]
result = create_grammar_backend(args, tokenizer, 32000, think_end_id=42)
self.assertIsInstance(result, ReasonerGrammarBackend)
- self.assertEqual(result.think_end_id, 42)
self.assertIs(result.grammar_backend, mock_backend)
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
@@ -396,7 +406,7 @@ class TestCreateGrammarBackend(unittest.TestCase):
"""Without think_end_id passed in, no reasoner wrapping."""
mock_backend = MagicMock(spec=BaseGrammarBackend)
mock_outlines_cls.return_value = mock_backend
- args = self._make_server_args("outlines", reasoning_parser="deepseek")
+ args = self._make_server_args("outlines", reasoning_parser="deepseek-r1")
tokenizer = MagicMock(spec=[]) # No think_end_id attribute
result = create_grammar_backend(args, tokenizer, 32000, think_end_id=None)
diff --git a/test/registered/unit/constrained/test_e2e_constrained_reasoning.py b/test/registered/unit/constrained/test_e2e_constrained_reasoning.py
new file mode 100644
index 000000000..48e69a4ac
--- /dev/null
+++ b/test/registered/unit/constrained/test_e2e_constrained_reasoning.py
@@ -0,0 +1,314 @@
+"""
+End-to-end tests for strict reasoning + constrained decoding.
+
+Tests that the full pipeline works:
+- AC-5.1: Strict reasoning + JSON schema constrained generation
+- AC-5.2: Strict reasoning + tool call parsing (basic validation only)
+
+These tests launch a real server with a small model and verify
+the constrained decoding pipeline produces valid output.
+"""
+
+import json
+import unittest
+
+import requests
+
+from sglang.srt.utils import kill_process_tree
+from sglang.test.ci.ci_register import register_cuda_ci
+from sglang.test.test_utils import (
+ DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
+ CustomTestCase,
+ popen_launch_server,
+)
+
+register_cuda_ci(est_time=120, suite="stage-b-test-1-gpu-small")
+
+MODEL = "Qwen/Qwen3-0.6B"
+BASE_URL = "http://127.0.0.1:39877"
+API_KEY = "sk-test-1234"
+
+
+class TestConstrainedReasoningE2E(CustomTestCase):
+ @classmethod
+ def setUpClass(cls):
+ cls.model = MODEL
+ cls.base_url = BASE_URL
+ cls.api_key = API_KEY
+ cls.process = popen_launch_server(
+ cls.model,
+ cls.base_url,
+ timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
+ api_key=cls.api_key,
+ other_args=[
+ "--reasoning-parser",
+ "qwen3",
+ ],
+ )
+
+ @classmethod
+ def tearDownClass(cls):
+ kill_process_tree(cls.process.pid)
+
+ def _chat(self, **kwargs):
+ default = {
+ "model": self.model,
+ "messages": [
+ {
+ "role": "user",
+ "content": "What is 2+2? Answer with just the number.",
+ }
+ ],
+ "temperature": 0,
+ "max_tokens": 256,
+ }
+ default.update(kwargs)
+ resp = requests.post(
+ f"{self.base_url}/v1/chat/completions",
+ headers={"Authorization": f"Bearer {self.api_key}"},
+ json=default,
+ timeout=60,
+ )
+ self.assertEqual(resp.status_code, 200, f"Request failed: {resp.text}")
+ return resp.json()
+
+ def test_reasoning_with_json_schema(self):
+ """AC-5.1: Reasoning + JSON schema produces valid JSON output."""
+ schema = {
+ "type": "object",
+ "properties": {
+ "answer": {"type": "integer"},
+ },
+ "required": ["answer"],
+ }
+ data = self._chat(
+ response_format={
+ "type": "json_schema",
+ "json_schema": {
+ "name": "answer_schema",
+ "schema": schema,
+ },
+ },
+ chat_template_kwargs={"enable_thinking": True},
+ separate_reasoning=True,
+ )
+
+ choice = data["choices"][0]
+ content = choice["message"]["content"] or ""
+
+ # Content should be valid JSON conforming to schema when non-empty.
+ # With small models + separate_reasoning, content may be empty if the
+ # model puts everything in reasoning_content. That's acceptable.
+ if content.strip():
+ try:
+ parsed = json.loads(content)
+ self.assertIn("answer", parsed)
+ self.assertIsInstance(parsed["answer"], int)
+ except (json.JSONDecodeError, TypeError):
+ # Small models may produce imperfect JSON
+ self.assertTrue(
+ content.strip().startswith("{"),
+ f"Expected JSON-like output, got: {content!r}",
+ )
+
+ # Content should NOT contain tags (those go to reasoning_content)
+ self.assertNotIn("", content)
+
+ def test_reasoning_disabled_with_json_schema(self):
+ """JSON schema still works when reasoning is explicitly disabled."""
+ schema = {
+ "type": "object",
+ "properties": {
+ "answer": {"type": "integer"},
+ },
+ "required": ["answer"],
+ }
+ data = self._chat(
+ response_format={
+ "type": "json_schema",
+ "json_schema": {
+ "name": "answer_schema",
+ "schema": schema,
+ },
+ },
+ chat_template_kwargs={"enable_thinking": False},
+ )
+
+ choice = data["choices"][0]
+ content = choice["message"]["content"]
+
+ # Should still produce valid JSON
+ parsed = json.loads(content)
+ self.assertIn("answer", parsed)
+
+ def test_reasoning_with_separate_output(self):
+ """Reasoning content is correctly separated from normal content."""
+ data = self._chat(
+ chat_template_kwargs={"enable_thinking": True},
+ separate_reasoning=True,
+ )
+
+ choice = data["choices"][0]
+ content = choice["message"]["content"]
+ reasoning = choice["message"].get("reasoning_content")
+
+ # Content should not contain think tags
+ self.assertNotIn("", content)
+ self.assertNotIn("", content)
+
+ def test_tool_call_after_reasoning(self):
+ """AC-5.2: Tool call parsing works with reasoning enabled."""
+ tools = [
+ {
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "description": "Get the current weather",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {"type": "string"},
+ },
+ "required": ["location"],
+ },
+ },
+ }
+ ]
+ data = self._chat(
+ messages=[
+ {
+ "role": "user",
+ "content": "What's the weather in Paris?",
+ }
+ ],
+ tools=tools,
+ chat_template_kwargs={"enable_thinking": True},
+ separate_reasoning=True,
+ )
+
+ choice = data["choices"][0]
+ # The model may or may not produce tool calls (depends on model capability)
+ # but the response should be well-formed (no crashes)
+ self.assertIn("message", choice)
+ self.assertIn("finish_reason", choice)
+ # finish_reason should be either "stop" or "tool_calls"
+ self.assertIn(choice["finish_reason"], ["stop", "tool_calls", "length"])
+
+
+class TestStrictThinkingE2E(CustomTestCase):
+ """E2E tests with --enable-strict-thinking flag.
+
+ Validates that the strict thinking flag is correctly propagated through
+ the full pipeline: server_args -> grammar_backend -> ReasonerGrammarBackend
+ -> token filtering during thinking phase.
+ """
+
+ @classmethod
+ def setUpClass(cls):
+ cls.model = MODEL
+ cls.base_url = "http://127.0.0.1:39878"
+ cls.api_key = API_KEY
+ cls.process = popen_launch_server(
+ cls.model,
+ cls.base_url,
+ timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
+ api_key=cls.api_key,
+ other_args=[
+ "--reasoning-parser",
+ "qwen3",
+ "--enable-strict-thinking",
+ ],
+ )
+
+ @classmethod
+ def tearDownClass(cls):
+ kill_process_tree(cls.process.pid)
+
+ def _chat(self, **kwargs):
+ default = {
+ "model": self.model,
+ "messages": [
+ {
+ "role": "user",
+ "content": "What is 2+2? Answer with just the number.",
+ }
+ ],
+ "temperature": 0,
+ "max_tokens": 256,
+ }
+ default.update(kwargs)
+ resp = requests.post(
+ f"{self.base_url}/v1/chat/completions",
+ headers={"Authorization": f"Bearer {self.api_key}"},
+ json=default,
+ timeout=60,
+ )
+ self.assertEqual(resp.status_code, 200, f"Request failed: {resp.text}")
+ return resp.json()
+
+ def test_strict_thinking_with_json_schema(self):
+ """Strict thinking + JSON schema: server starts and produces valid output."""
+ schema = {
+ "type": "object",
+ "properties": {
+ "answer": {"type": "integer"},
+ },
+ "required": ["answer"],
+ }
+ data = self._chat(
+ response_format={
+ "type": "json_schema",
+ "json_schema": {
+ "name": "answer_schema",
+ "schema": schema,
+ },
+ },
+ chat_template_kwargs={"enable_thinking": True},
+ separate_reasoning=True,
+ )
+
+ choice = data["choices"][0]
+ content = choice["message"]["content"] or ""
+
+ if content.strip():
+ try:
+ parsed = json.loads(content)
+ self.assertIn("answer", parsed)
+ except (json.JSONDecodeError, TypeError):
+ self.assertTrue(
+ content.strip().startswith("{"),
+ f"Expected JSON-like output, got: {content!r}",
+ )
+
+ # Think tags must not leak into content
+ self.assertNotIn("", content)
+
+ def test_strict_thinking_disabled_per_request(self):
+ """When thinking is disabled per-request, strict server still works."""
+ data = self._chat(
+ chat_template_kwargs={"enable_thinking": False},
+ )
+
+ choice = data["choices"][0]
+ self.assertIn("message", choice)
+ self.assertIn("finish_reason", choice)
+ # Should complete normally without errors
+ self.assertIn(choice["finish_reason"], ["stop", "length"])
+
+ def test_strict_thinking_separate_reasoning(self):
+ """Strict thinking with separate_reasoning produces well-formed output."""
+ data = self._chat(
+ chat_template_kwargs={"enable_thinking": True},
+ separate_reasoning=True,
+ )
+
+ choice = data["choices"][0]
+ content = choice["message"]["content"] or ""
+
+ # Think tags must not leak into content
+ self.assertNotIn("", content)
+ self.assertNotIn("", content)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/test/registered/unit/constrained/test_grammar_manager.py b/test/registered/unit/constrained/test_grammar_manager.py
index 649bf49f4..2d25a4bc4 100644
--- a/test/registered/unit/constrained/test_grammar_manager.py
+++ b/test/registered/unit/constrained/test_grammar_manager.py
@@ -24,6 +24,7 @@ from sglang.srt.constrained.base_grammar_backend import (
InvalidGrammarObject,
)
from sglang.srt.constrained.grammar_manager import GrammarManager
+from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(2.0, "stage-a-test-cpu")
@@ -48,7 +49,12 @@ def _make_scheduler(grammar_backend_name="none", skip_tokenizer=False):
def _make_req(
- json_schema=None, regex=None, ebnf=None, structural_tag=None, rid="req-1"
+ json_schema=None,
+ regex=None,
+ ebnf=None,
+ structural_tag=None,
+ rid="req-1",
+ custom_params=None,
):
"""Create a mock request with sampling params."""
req = MagicMock()
@@ -57,6 +63,7 @@ def _make_req(
req.sampling_params.regex = regex
req.sampling_params.ebnf = ebnf
req.sampling_params.structural_tag = structural_tag
+ req.sampling_params.custom_params = custom_params
req.require_reasoning = False
req.grammar = None
req.grammar_key = None
@@ -256,6 +263,39 @@ class TestProcessReqWithGrammar(unittest.TestCase):
self.assertTrue(mgr.has_waiting_grammars())
self.assertEqual(len(mgr), 1)
+ def test_cache_hit_applies_request_thinking_budget(self):
+ mgr = self._make_mgr()
+ grammar_obj = ReasonerGrammarObject(
+ grammar=None, think_end_id=0, max_think_tokens=99
+ )
+ mgr.grammar_backend.get_cached_or_future_value.return_value = (
+ grammar_obj,
+ True,
+ )
+
+ req = _make_req(
+ json_schema="schema",
+ custom_params={"thinking_budget": 7},
+ )
+ mgr.process_req_with_grammar(req)
+
+ self.assertEqual(req.grammar.max_think_tokens, 7)
+
+ def test_strict_reasoning_grammar_applies_request_thinking_budget(self):
+ mgr = self._make_mgr()
+ mgr._enable_strict_thinking = True
+ grammar_obj = ReasonerGrammarObject(
+ grammar=None, think_end_id=0, max_think_tokens=99
+ )
+ mgr.grammar_backend.init_strict_reasoning_grammar.return_value = grammar_obj
+
+ req = _make_req(custom_params={"thinking_budget": 3})
+ req.require_reasoning = True
+ mgr.process_req_with_grammar(req)
+
+ self.assertIs(req.grammar, grammar_obj)
+ self.assertEqual(req.grammar.max_think_tokens, 3)
+
class TestAbortRequests(unittest.TestCase):
"""Test abort_requests handling."""
@@ -494,8 +534,8 @@ class TestGetReadyGrammarRequests(unittest.TestCase):
req.set_finish_with_abort.assert_called_once()
self.assertIn("timed out", req.set_finish_with_abort.call_args[0][0])
- def test_future_exception_propagates(self):
- """A future that raised an exception should propagate on .result()."""
+ def test_future_exception_creates_invalid_grammar_object(self):
+ """A future that raised an exception should create InvalidGrammarObject, not crash."""
mgr = self._make_mgr()
future = Future()
@@ -506,8 +546,32 @@ class TestGetReadyGrammarRequests(unittest.TestCase):
req.grammar_key = ("json", "crash")
mgr.grammar_queue.append(req)
- with self.assertRaises(RuntimeError):
- mgr.get_ready_grammar_requests()
+ result = mgr.get_ready_grammar_requests()
+ self.assertEqual(len(result), 1)
+ self.assertIsInstance(result[0].grammar, InvalidGrammarObject)
+ req.set_finish_with_abort.assert_called_once()
+
+ def test_ready_future_applies_request_budget_without_polluting_cache(self):
+ mgr = self._make_mgr()
+
+ grammar_obj = ReasonerGrammarObject(
+ grammar=None, think_end_id=0, max_think_tokens=99
+ )
+ future = Future()
+ future.set_result(grammar_obj)
+
+ req = _make_req(json_schema="schema", custom_params={"thinking_budget": 4})
+ req.grammar = future
+ req.grammar_key = ("json", "schema")
+ mgr.grammar_queue.append(req)
+
+ result = mgr.get_ready_grammar_requests()
+
+ self.assertEqual(len(result), 1)
+ self.assertEqual(req.grammar.max_think_tokens, 4)
+ cached_key, cached_value = mgr.grammar_backend.set_cache.call_args[0]
+ self.assertEqual(cached_key, ("json", "schema"))
+ self.assertEqual(cached_value.max_think_tokens, 99)
@patch("sglang.srt.constrained.grammar_manager.torch.distributed.all_gather_object")
def test_multi_rank_sync_intersects_ready_unions_failed(self, mock_all_gather):
@@ -579,5 +643,100 @@ class TestGetReadyGrammarRequests(unittest.TestCase):
self.assertEqual(len(mgr.grammar_queue), 0)
+class TestStrictReasoningPaths(unittest.TestCase):
+ """Test _enable_strict_thinking code paths in GrammarManager."""
+
+ def _make_mgr(self):
+ scheduler = _make_scheduler()
+ scheduler.server_args.skip_tokenizer_init = True
+ mgr = GrammarManager(scheduler)
+ mgr.grammar_backend = MagicMock(spec=BaseGrammarBackend)
+ mgr._enable_strict_thinking = True
+ return mgr
+
+ def test_strict_unconstrained_request_gets_strict_grammar(self):
+ """Request without json_schema/regex/ebnf should get strict-only grammar."""
+ mgr = self._make_mgr()
+ grammar_obj = MagicMock()
+ mgr.grammar_backend.init_strict_reasoning_grammar.return_value = grammar_obj
+
+ req = _make_req() # No constraint
+ req.require_reasoning = True
+ result = mgr.process_req_with_grammar(req)
+
+ self.assertFalse(result) # Not added to grammar queue
+ self.assertIs(req.grammar, grammar_obj)
+ mgr.grammar_backend.init_strict_reasoning_grammar.assert_called_once_with(True)
+
+ def test_strict_unconstrained_no_reasoning_flag(self):
+ """Unconstrained request with require_reasoning=False still gets strict grammar."""
+ mgr = self._make_mgr()
+ grammar_obj = MagicMock()
+ mgr.grammar_backend.init_strict_reasoning_grammar.return_value = grammar_obj
+
+ req = _make_req()
+ req.require_reasoning = False
+ mgr.process_req_with_grammar(req)
+
+ self.assertIs(req.grammar, grammar_obj)
+ mgr.grammar_backend.init_strict_reasoning_grammar.assert_called_once_with(False)
+
+ def test_strict_unconstrained_none_grammar_is_fine(self):
+ """If init_strict_reasoning_grammar returns None, req.grammar stays None."""
+ mgr = self._make_mgr()
+ mgr.grammar_backend.init_strict_reasoning_grammar.return_value = None
+
+ req = _make_req()
+ req.require_reasoning = True
+ mgr.process_req_with_grammar(req)
+
+ self.assertIsNone(req.grammar)
+
+ def test_strict_constrained_request_uses_normal_dispatch(self):
+ """Request with json_schema should go through normal dispatch, not strict path."""
+ mgr = self._make_mgr()
+ future = MagicMock(spec=Future)
+ mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False)
+
+ req = _make_req(json_schema='{"type": "object"}')
+ req.require_reasoning = True
+ result = mgr.process_req_with_grammar(req)
+
+ self.assertTrue(result) # Added to grammar queue
+ mgr.grammar_backend.init_strict_reasoning_grammar.assert_not_called()
+
+ def test_strict_not_set_skips_strict_path(self):
+ """When _enable_strict_thinking=False, unconstrained requests get no grammar."""
+ mgr = self._make_mgr()
+ mgr._enable_strict_thinking = False
+
+ req = _make_req()
+ req.require_reasoning = True
+ mgr.process_req_with_grammar(req)
+
+ self.assertIsNone(req.grammar)
+ mgr.grammar_backend.init_strict_reasoning_grammar.assert_not_called()
+
+ def test_future_exception_creates_invalid_grammar(self):
+ """Future.result() raising should create InvalidGrammarObject, not crash."""
+ mgr = self._make_mgr()
+
+ future = Future()
+ future.set_exception(RuntimeError("compilation failed"))
+
+ req = _make_req(json_schema='{"type": "object"}')
+ req.require_reasoning = True
+ req.grammar = future
+ req.grammar_key = ("json", '{"type": "object"}')
+ mgr.grammar_queue.append(req)
+
+ mgr.SGLANG_GRAMMAR_POLL_INTERVAL = 0.001
+ result = mgr.get_ready_grammar_requests()
+
+ self.assertEqual(len(result), 1)
+ self.assertIsInstance(result[0].grammar, InvalidGrammarObject)
+ req.set_finish_with_abort.assert_called_once()
+
+
if __name__ == "__main__":
unittest.main()
diff --git a/test/registered/unit/constrained/test_reasoner_grammar_backend.py b/test/registered/unit/constrained/test_reasoner_grammar_backend.py
index 1f9afdb69..365d42c3c 100644
--- a/test/registered/unit/constrained/test_reasoner_grammar_backend.py
+++ b/test/registered/unit/constrained/test_reasoner_grammar_backend.py
@@ -1,413 +1,453 @@
-"""
-Unit tests for sglang.srt.constrained.reasoner_grammar_backend.
-
-Test Coverage:
-- ReasonerGrammarObject: state transitions, accept_token during thinking
- vs post-thinking, rollback across think boundary, fill_vocab_mask gating,
- copy semantics, finished delegation, delegation of jump methods
-- ReasonerGrammarBackend: dispatch wrapping, invalid grammar passthrough,
- None grammar passthrough, reasoning init on wrapped object
-
-Usage:
- python -m pytest test_reasoner_grammar_backend.py -v
-"""
-
+import os
import unittest
-from unittest.mock import MagicMock, call
+from types import SimpleNamespace
+from unittest.mock import MagicMock
-from sglang.srt.constrained.base_grammar_backend import (
- BaseGrammarBackend,
- BaseGrammarObject,
- InvalidGrammarObject,
-)
+import torch
+
+from sglang.srt.constrained.base_grammar_backend import BaseGrammarBackend
from sglang.srt.constrained.reasoner_grammar_backend import (
ReasonerGrammarBackend,
ReasonerGrammarObject,
)
+from sglang.srt.constrained.torch_ops.token_filter_torch_ops import (
+ set_token_filter_torch,
+)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(2.0, "stage-a-test-cpu")
-THINK_END_ID = 99
+
+class _DummyTokenizer:
+ def __init__(self, token_map):
+ self._token_map = token_map
+
+ def encode(self, text, add_special_tokens=False):
+ return list(self._token_map.get(text, []))
-class TestReasonerGrammarObjectStateTransitions(unittest.TestCase):
- """Test thinking state machine in ReasonerGrammarObject."""
+class _DummyGrammarBackend(BaseGrammarBackend):
+ def __init__(self, support_token_filter=True):
+ super().__init__()
+ self._support_token_filter = support_token_filter
+ self._dispatch_result = None
- def _make(self):
- grammar = MagicMock(spec=BaseGrammarObject)
- return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
+ @property
+ def is_support_token_filter(self):
+ return self._support_token_filter
- def test_initial_state_thinking(self):
- obj, _ = self._make()
- self.assertEqual(obj.tokens_after_think_end, -1)
+ @staticmethod
+ def allocate_vocab_mask(vocab_size, batch_size, device):
+ return torch.zeros((batch_size, (vocab_size + 31) // 32), dtype=torch.int32)
- def test_transfer_state_during_thinking(self):
- """Regular tokens during thinking don't change state."""
- obj, _ = self._make()
- obj.transfer_state(10)
- self.assertEqual(obj.tokens_after_think_end, -1)
+ @staticmethod
+ def move_vocab_mask(vocab_mask, device):
+ return vocab_mask
- def test_transfer_state_think_end_token(self):
- """Think end token transitions from -1 to 0."""
- obj, _ = self._make()
- obj.transfer_state(THINK_END_ID)
- self.assertEqual(obj.tokens_after_think_end, 0)
+ @staticmethod
+ def apply_vocab_mask(logits, vocab_mask):
+ return None
- def test_transfer_state_increments_after_thinking(self):
- """After thinking ends, each token increments counter."""
- obj, _ = self._make()
- obj.tokens_after_think_end = 0
- obj.transfer_state(10)
- self.assertEqual(obj.tokens_after_think_end, 1)
- obj.transfer_state(20)
- self.assertEqual(obj.tokens_after_think_end, 2)
+ @staticmethod
+ def set_token_filter(
+ vocab_mask, token_ids, batch_idx, is_allowed=True, reset_vocab_mask=True
+ ):
+ set_token_filter_torch(
+ vocab_mask, token_ids, batch_idx, is_allowed, reset_vocab_mask
+ )
- def test_think_end_after_thinking_already_ended(self):
- """Second think_end_id after thinking ended just increments."""
- obj, _ = self._make()
- obj.tokens_after_think_end = 3
- obj.transfer_state(THINK_END_ID)
- self.assertEqual(obj.tokens_after_think_end, 4)
-
- def test_rollback_state_from_post_thinking(self):
- obj, _ = self._make()
- obj.tokens_after_think_end = 3
- obj.rollback_state()
- self.assertEqual(obj.tokens_after_think_end, 2)
-
- def test_rollback_state_at_boundary(self):
- """Rollback from 0 goes back to -1 (thinking)."""
- obj, _ = self._make()
- obj.tokens_after_think_end = 0
- obj.rollback_state()
- self.assertEqual(obj.tokens_after_think_end, -1)
-
- def test_rollback_state_during_thinking(self):
- """Rollback during thinking stays at -1."""
- obj, _ = self._make()
- obj.rollback_state()
- self.assertEqual(obj.tokens_after_think_end, -1)
+ def _init_value_dispatch(self, key, reasoning):
+ return self._dispatch_result
-class TestReasonerGrammarObjectAcceptToken(unittest.TestCase):
- """Test accept_token behavior with thinking/post-thinking states."""
-
- def _make(self):
- grammar = MagicMock(spec=BaseGrammarObject)
- return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
-
- def test_accept_during_thinking_skips_grammar(self):
- """During thinking phase, inner grammar should NOT receive tokens."""
- obj, grammar = self._make()
- obj.accept_token(10)
- grammar.accept_token.assert_not_called()
- # State should still be -1
- self.assertEqual(obj.tokens_after_think_end, -1)
-
- def test_accept_think_end_token(self):
- """Think end token transitions state but doesn't call inner grammar (state was -1 before transfer)."""
- obj, grammar = self._make()
- # tokens_after_think_end is -1, so grammar.accept_token is not called
- # But wait: accept_token checks `>= 0` BEFORE transfer_state
- # At call time tokens_after_think_end == -1, so grammar.accept_token skipped
- obj.accept_token(THINK_END_ID)
- grammar.accept_token.assert_not_called()
- self.assertEqual(obj.tokens_after_think_end, 0)
-
- def test_accept_after_thinking_calls_grammar(self):
- """After thinking ends, tokens go to inner grammar."""
- obj, grammar = self._make()
- obj.tokens_after_think_end = 0
- obj.accept_token(42)
- grammar.accept_token.assert_called_once_with(42)
- self.assertEqual(obj.tokens_after_think_end, 1)
-
- def test_accept_sequence_through_thinking_and_generation(self):
- """Full sequence: think tokens -> think_end -> generation tokens."""
- obj, grammar = self._make()
-
- # Thinking phase
- obj.accept_token(1)
- obj.accept_token(2)
- self.assertEqual(grammar.accept_token.call_count, 0)
-
- # Think end
- obj.accept_token(THINK_END_ID)
- self.assertEqual(grammar.accept_token.call_count, 0)
-
- # Generation phase
- obj.accept_token(10)
- obj.accept_token(20)
- self.assertEqual(grammar.accept_token.call_count, 2)
- grammar.accept_token.assert_has_calls([call(10), call(20)])
+def _allowed_token_ids(vocab_mask, token_ids):
+ allowed = []
+ for token_id in token_ids:
+ elem = token_id // 32
+ bit = token_id % 32
+ if int(vocab_mask[0, elem].item()) & (1 << bit):
+ allowed.append(token_id)
+ return allowed
-class TestReasonerGrammarObjectRollback(unittest.TestCase):
- """Test rollback across thinking boundary."""
+class TestReasonerGrammarObject(unittest.TestCase):
+ def _make_strict_object(self):
+ return ReasonerGrammarObject(
+ grammar=None,
+ think_end_id=7,
+ think_excluded_token_ids=[3, 5],
+ max_think_tokens=2,
+ enable_token_filter=True,
+ token_filter_fn=set_token_filter_torch,
+ allocate_vocab_mask_fn=lambda vocab_size, batch_size, device: torch.zeros(
+ (batch_size, (vocab_size + 31) // 32), dtype=torch.int32
+ ),
+ move_vocab_mask_fn=lambda vocab_mask, device: vocab_mask,
+ apply_vocab_mask_fn=lambda logits, vocab_mask: None,
+ )
- def _make(self):
- grammar = MagicMock(spec=BaseGrammarObject)
- return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
-
- def test_rollback_within_generation(self):
- """Rollback entirely within generation phase."""
- obj, grammar = self._make()
- obj.tokens_after_think_end = 5
- obj.rollback(3)
- grammar.rollback.assert_called_once_with(3)
- self.assertEqual(obj.tokens_after_think_end, 2)
-
- def test_rollback_across_boundary(self):
- """Rollback that crosses from generation back into thinking."""
- obj, grammar = self._make()
- obj.tokens_after_think_end = 2
- obj.rollback(4)
- # Only 2 tokens were post-thinking, so inner grammar rolls back 2
- grammar.rollback.assert_called_once_with(2)
- # After 4 rollback_state calls from 2: 2->1->0->-1->-1
- self.assertEqual(obj.tokens_after_think_end, -1)
-
- def test_rollback_during_thinking(self):
- """Rollback during thinking phase doesn't touch inner grammar."""
- obj, grammar = self._make()
- obj.rollback(3)
- grammar.rollback.assert_not_called()
- self.assertEqual(obj.tokens_after_think_end, -1)
-
- def test_rollback_zero(self):
- obj, grammar = self._make()
- obj.tokens_after_think_end = 2
- obj.rollback(0)
- grammar.rollback.assert_not_called()
- self.assertEqual(obj.tokens_after_think_end, 2)
-
- def test_rollback_exactly_to_boundary(self):
- """Rollback exactly the number of post-thinking tokens."""
- obj, grammar = self._make()
- obj.tokens_after_think_end = 3
- obj.rollback(3)
- grammar.rollback.assert_called_once_with(3)
- self.assertEqual(obj.tokens_after_think_end, 0)
-
- def test_rollback_far_beyond_all_tokens(self):
- """Rollback k much larger than tokens_after_think_end clamps grammar rollback."""
- obj, grammar = self._make()
- obj.tokens_after_think_end = 2
- obj.rollback(100)
- # Inner grammar only rolls back the 2 post-thinking tokens
- grammar.rollback.assert_called_once_with(2)
- # State bottoms out at -1
- self.assertEqual(obj.tokens_after_think_end, -1)
-
- def test_accept_then_rollback_roundtrip(self):
- """Accept tokens then rollback should restore original state."""
- obj, grammar = self._make()
- obj.tokens_after_think_end = 0 # Just finished thinking
-
- # Accept 3 generation tokens
- obj.accept_token(10)
- obj.accept_token(20)
- obj.accept_token(30)
- self.assertEqual(obj.tokens_after_think_end, 3)
- self.assertEqual(grammar.accept_token.call_count, 3)
-
- # Rollback all 3
- obj.rollback(3)
- self.assertEqual(obj.tokens_after_think_end, 0)
- grammar.rollback.assert_called_once_with(3)
-
-
-class TestReasonerGrammarObjectVocabMask(unittest.TestCase):
- """Test vocab mask gating based on thinking state."""
-
- def _make(self):
- grammar = MagicMock(spec=BaseGrammarObject)
- return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
-
- def test_fill_during_thinking_skips(self):
- obj, grammar = self._make()
- obj.fill_vocab_mask("mask", 0)
- grammar.fill_vocab_mask.assert_not_called()
-
- def test_fill_after_thinking_delegates(self):
- obj, grammar = self._make()
- obj.tokens_after_think_end = 0
- obj.fill_vocab_mask("mask", 0)
- grammar.fill_vocab_mask.assert_called_once_with("mask", 0)
-
- def test_fill_well_into_generation(self):
- obj, grammar = self._make()
- obj.tokens_after_think_end = 5
- obj.fill_vocab_mask("mask", 2)
- grammar.fill_vocab_mask.assert_called_once_with("mask", 2)
-
- def test_fill_at_think_end_boundary(self):
- """After accepting think_end token, fill_vocab_mask should delegate."""
- obj, grammar = self._make()
- # Simulate: accept think_end, state goes from -1 to 0
- obj.accept_token(THINK_END_ID)
- self.assertEqual(obj.tokens_after_think_end, 0)
- obj.fill_vocab_mask("mask", 0)
- grammar.fill_vocab_mask.assert_called_once_with("mask", 0)
-
- def test_allocate_delegates(self):
- obj, grammar = self._make()
- obj.allocate_vocab_mask(32000, 4, "cpu")
- grammar.allocate_vocab_mask.assert_called_once_with(32000, 4, "cpu")
-
- def test_move_delegates(self):
- obj, grammar = self._make()
- obj.move_vocab_mask("mask", "cuda")
- grammar.move_vocab_mask.assert_called_once_with("mask", "cuda")
-
-
-class TestReasonerGrammarObjectDelegation(unittest.TestCase):
- """Test that non-state methods delegate to inner grammar."""
-
- def _make(self):
- grammar = MagicMock(spec=BaseGrammarObject)
- return ReasonerGrammarObject(grammar, THINK_END_ID), grammar
-
- def test_is_terminated_delegates(self):
- obj, grammar = self._make()
- grammar.is_terminated.return_value = True
- self.assertTrue(obj.is_terminated())
-
- def test_finished_getter_delegates(self):
- obj, grammar = self._make()
- grammar.finished = True
- self.assertTrue(obj.finished)
-
- def test_finished_setter_delegates(self):
- obj, grammar = self._make()
- obj.finished = True
- self.assertTrue(grammar.finished)
-
- def test_try_jump_forward_delegates(self):
- obj, grammar = self._make()
- grammar.try_jump_forward.return_value = ([1, 2], "ab")
- result = obj.try_jump_forward("tokenizer")
- grammar.try_jump_forward.assert_called_once_with("tokenizer")
- self.assertEqual(result, ([1, 2], "ab"))
-
- def test_jump_forward_str_state_delegates(self):
- obj, grammar = self._make()
- grammar.jump_forward_str_state.return_value = ("str", 5)
- result = obj.jump_forward_str_state("helper")
- self.assertEqual(result, ("str", 5))
-
- def test_jump_and_retokenize_delegates(self):
- obj, grammar = self._make()
- obj.jump_and_retokenize([1], [2], 3)
- grammar.jump_and_retokenize.assert_called_once_with([1], [2], 3)
-
- def test_apply_vocab_mask_property(self):
- obj, grammar = self._make()
- grammar.apply_vocab_mask = "mask_fn"
- self.assertEqual(obj.apply_vocab_mask, "mask_fn")
-
- def test_copy_creates_new_wrapper(self):
- obj, grammar = self._make()
- grammar_copy = MagicMock(spec=BaseGrammarObject)
- grammar.copy.return_value = grammar_copy
-
- copied = obj.copy()
- self.assertIsInstance(copied, ReasonerGrammarObject)
- self.assertIsNot(copied, obj)
- self.assertIs(copied.grammar, grammar_copy)
- self.assertEqual(copied.think_end_id, THINK_END_ID)
-
- def test_copy_does_not_share_state(self):
- """Modifying copy's state should not affect the original."""
- obj, grammar = self._make()
- grammar_copy = MagicMock(spec=BaseGrammarObject)
- grammar.copy.return_value = grammar_copy
-
- copied = obj.copy()
- copied.tokens_after_think_end = 5
- self.assertEqual(obj.tokens_after_think_end, -1)
-
-
-class TestReasonerGrammarObjectMaybeInitReasoning(unittest.TestCase):
- """Test maybe_init_reasoning state initialization."""
-
- def test_reasoning_true_sets_thinking(self):
- grammar = MagicMock(spec=BaseGrammarObject)
- obj = ReasonerGrammarObject(grammar, THINK_END_ID)
+ def test_strict_thinking_phase_excludes_configured_tokens(self):
+ obj = self._make_strict_object()
obj.maybe_init_reasoning(True)
- self.assertEqual(obj.tokens_after_think_end, -1)
+ mask = obj.allocate_vocab_mask(64, 1, "cpu")
- def test_reasoning_false_skips_thinking(self):
- grammar = MagicMock(spec=BaseGrammarObject)
- obj = ReasonerGrammarObject(grammar, THINK_END_ID)
- obj.maybe_init_reasoning(False)
- self.assertEqual(obj.tokens_after_think_end, 0)
+ obj.fill_vocab_mask(mask, 0)
- def test_reasoning_toggle(self):
- """Toggling reasoning resets state regardless of current position."""
- grammar = MagicMock(spec=BaseGrammarObject)
- obj = ReasonerGrammarObject(grammar, THINK_END_ID)
- obj.tokens_after_think_end = 5 # Deep into generation
+ allowed = _allowed_token_ids(mask, [0, 1, 3, 5, 7, 8])
+ self.assertEqual(allowed, [0, 1, 7, 8])
+ def test_budget_exhaustion_allows_only_think_end(self):
+ obj = self._make_strict_object()
obj.maybe_init_reasoning(True)
- self.assertEqual(obj.tokens_after_think_end, -1)
+ obj.accept_token(10)
+ obj.accept_token(11)
+ mask = obj.allocate_vocab_mask(64, 1, "cpu")
- obj.maybe_init_reasoning(False)
- self.assertEqual(obj.tokens_after_think_end, 0)
+ obj.fill_vocab_mask(mask, 0)
+
+ allowed = _allowed_token_ids(mask, [0, 1, 3, 5, 7, 8, 10, 11])
+ self.assertEqual(allowed, [7])
+
+ def test_strict_only_wrapper_exposes_backend_mask_hooks(self):
+ obj = self._make_strict_object()
+ mask = obj.allocate_vocab_mask(64, 2, "cpu")
+
+ self.assertEqual(mask.shape, (2, 2))
+ self.assertIs(obj.move_vocab_mask(mask, "cpu"), mask)
+ self.assertIsNotNone(obj.apply_vocab_mask)
class TestReasonerGrammarBackend(unittest.TestCase):
- """Test ReasonerGrammarBackend dispatch wrapping."""
+ def setUp(self):
+ self._prev_budget = os.environ.get("SGLANG_MAX_THINK_TOKENS")
- def _make(self):
- inner = MagicMock(spec=BaseGrammarBackend)
- backend = ReasonerGrammarBackend(inner, THINK_END_ID)
- return backend, inner
+ def tearDown(self):
+ if self._prev_budget is None:
+ os.environ.pop("SGLANG_MAX_THINK_TOKENS", None)
+ else:
+ os.environ["SGLANG_MAX_THINK_TOKENS"] = self._prev_budget
- def test_wraps_valid_grammar(self):
- backend, inner = self._make()
- mock_grammar = MagicMock(spec=BaseGrammarObject)
- inner._init_value_dispatch.return_value = mock_grammar
+ def _make_parser(self):
+ detector = SimpleNamespace(
+ think_start_token="",
+ think_end_token="",
+ think_excluded_tokens=["", ""],
+ )
+ return SimpleNamespace(detector=detector)
- result = backend._init_value_dispatch(("json", "schema"), True)
- self.assertIsInstance(result, ReasonerGrammarObject)
- self.assertIs(result.grammar, mock_grammar)
- self.assertEqual(result.think_end_id, THINK_END_ID)
+ def _make_tokenizer(self, start_ids=None, end_ids=None):
+ return _DummyTokenizer(
+ {
+ "": [1] if start_ids is None else start_ids,
+ "": [2] if end_ids is None else end_ids,
+ "": [3],
+ "": [4],
+ }
+ )
- def test_passes_through_invalid_grammar(self):
- backend, inner = self._make()
- invalid = InvalidGrammarObject("bad grammar")
- inner._init_value_dispatch.return_value = invalid
+ def test_init_strict_reasoning_grammar_uses_token_filter_and_budget(self):
+ os.environ["SGLANG_MAX_THINK_TOKENS"] = "2"
+ backend = _DummyGrammarBackend(support_token_filter=True)
+ reasoner = ReasonerGrammarBackend(
+ backend,
+ self._make_parser(),
+ self._make_tokenizer(),
+ enable_strict_thinking=True,
+ )
- result = backend._init_value_dispatch(("json", "schema"), False)
- self.assertIs(result, invalid)
- self.assertIsInstance(result, InvalidGrammarObject)
+ obj = reasoner.init_strict_reasoning_grammar(reasoning=True)
- def test_passes_through_none(self):
- backend, inner = self._make()
- inner._init_value_dispatch.return_value = None
+ self.assertIsInstance(obj, ReasonerGrammarObject)
+ self.assertTrue(obj.enable_token_filter)
+ self.assertEqual(obj.max_think_tokens, 2)
+ self.assertEqual(obj.think_excluded_token_ids, [3, 4])
- result = backend._init_value_dispatch(("json", "schema"), False)
- self.assertIsNone(result)
+ def test_init_strict_reasoning_grammar_none_when_strict_disabled(self):
+ backend = _DummyGrammarBackend(support_token_filter=True)
+ reasoner = ReasonerGrammarBackend(
+ backend,
+ self._make_parser(),
+ self._make_tokenizer(),
+ enable_strict_thinking=False,
+ )
- def test_inits_reasoning_on_wrapped(self):
- backend, inner = self._make()
- mock_grammar = MagicMock(spec=BaseGrammarObject)
- inner._init_value_dispatch.return_value = mock_grammar
+ self.assertIsNone(reasoner.init_strict_reasoning_grammar(reasoning=True))
- result = backend._init_value_dispatch(("json", "schema"), True)
- # reasoning=True → tokens_after_think_end should be -1
- self.assertEqual(result.tokens_after_think_end, -1)
+ def test_wraps_inner_grammar_with_reasoning_state_machine(self):
+ os.environ["SGLANG_MAX_THINK_TOKENS"] = "1"
+ backend = _DummyGrammarBackend(support_token_filter=True)
+ inner_grammar = MagicMock()
+ backend._dispatch_result = inner_grammar
+ reasoner = ReasonerGrammarBackend(
+ backend,
+ self._make_parser(),
+ self._make_tokenizer(),
+ enable_strict_thinking=True,
+ )
- def test_inits_no_reasoning_on_wrapped(self):
- backend, inner = self._make()
- mock_grammar = MagicMock(spec=BaseGrammarObject)
- inner._init_value_dispatch.return_value = mock_grammar
+ wrapped = reasoner._init_value_dispatch(("json", "{}"), reasoning=True)
+ self.assertIsInstance(wrapped, ReasonerGrammarObject)
+ wrapped.accept_token(10)
+ inner_grammar.accept_token.assert_not_called()
+ wrapped.accept_token(2)
+ wrapped.accept_token(42)
+ inner_grammar.accept_token.assert_called_once_with(42)
- result = backend._init_value_dispatch(("json", "schema"), False)
- # reasoning=False → tokens_after_think_end should be 0
- self.assertEqual(result.tokens_after_think_end, 0)
+ def test_accepts_multi_token_think_start_marker(self):
+ """think_start_token can be multi-token (e.g., GPT-OSS) since it's not used."""
+ backend = _DummyGrammarBackend(support_token_filter=True)
+ reasoner = ReasonerGrammarBackend(
+ backend,
+ self._make_parser(),
+ self._make_tokenizer(start_ids=[1, 2]),
+ enable_strict_thinking=True,
+ )
+ self.assertIsNotNone(reasoner)
+
+ def test_rejects_multi_token_think_end_marker(self):
+ backend = _DummyGrammarBackend(support_token_filter=True)
+
+ with self.assertRaisesRegex(ValueError, "must encode to exactly one token"):
+ ReasonerGrammarBackend(
+ backend,
+ self._make_parser(),
+ self._make_tokenizer(end_ids=[2, 3]),
+ enable_strict_thinking=True,
+ )
+
+ def test_rejects_unencodable_excluded_token(self):
+ backend = _DummyGrammarBackend(support_token_filter=True)
+ parser = self._make_parser()
+ parser.detector.think_excluded_tokens = [""]
+ tokenizer = _DummyTokenizer(
+ {
+ "": [1],
+ "": [2],
+ }
+ )
+
+ with self.assertRaisesRegex(ValueError, "could not be encoded"):
+ ReasonerGrammarBackend(
+ backend,
+ parser,
+ tokenizer,
+ enable_strict_thinking=True,
+ )
+
+ def test_strict_mode_fails_when_backend_lacks_token_filter(self):
+ backend = _DummyGrammarBackend(support_token_filter=False)
+
+ with self.assertRaisesRegex(ValueError, "does not support token filtering"):
+ ReasonerGrammarBackend(
+ backend,
+ self._make_parser(),
+ self._make_tokenizer(),
+ enable_strict_thinking=True,
+ )
+
+
+class TestReasonerGrammarObjectRollback(unittest.TestCase):
+ """Tests for rollback correctness at the THINKING→GENERATION boundary."""
+
+ def _make_object_with_mock_grammar(self):
+ inner_grammar = MagicMock()
+ inner_grammar.is_terminated.return_value = False
+ obj = ReasonerGrammarObject(
+ grammar=inner_grammar,
+ think_end_id=7,
+ think_excluded_token_ids=[3, 5],
+ max_think_tokens=-1,
+ enable_token_filter=True,
+ token_filter_fn=set_token_filter_torch,
+ allocate_vocab_mask_fn=lambda vs, bs, d: torch.zeros(
+ (bs, (vs + 31) // 32), dtype=torch.int32
+ ),
+ move_vocab_mask_fn=lambda vm, d: vm,
+ apply_vocab_mask_fn=lambda l, vm: None,
+ )
+ return obj, inner_grammar
+
+ def test_rollback_at_generation_boundary_returns_to_thinking(self):
+ obj, inner_grammar = self._make_object_with_mock_grammar()
+ obj.maybe_init_reasoning(True)
+
+ # Accept 3 thinking tokens then think_end_id
+ obj.accept_token(10)
+ obj.accept_token(11)
+ obj.accept_token(12)
+ obj.accept_token(7) # think_end_id → tokens_after_end = 0
+
+ self.assertTrue(obj._is_generation())
+ self.assertEqual(obj.tokens_after_end, 0)
+
+ # Rollback 1 step: should return to THINKING
+ obj.rollback(1)
+ self.assertTrue(obj._is_thinking())
+ self.assertEqual(obj.tokens_in_think, 3)
+ self.assertEqual(obj.tokens_after_end, -1)
+ # Grammar should not have been rolled back (no generation tokens were accepted)
+ inner_grammar.rollback.assert_not_called()
+
+ def test_rollback_spanning_both_phases(self):
+ obj, inner_grammar = self._make_object_with_mock_grammar()
+ obj.maybe_init_reasoning(True)
+
+ # 2 thinking tokens + think_end + 3 generation tokens
+ obj.accept_token(10) # think
+ obj.accept_token(11) # think
+ obj.accept_token(7) # think_end_id
+ obj.accept_token(20) # gen 1
+ obj.accept_token(21) # gen 2
+ obj.accept_token(22) # gen 3
+
+ self.assertEqual(obj.tokens_after_end, 3)
+
+ # Rollback 5: should roll back 3 generation tokens + think_end + 1 thinking token
+ obj.rollback(5)
+ self.assertTrue(obj._is_thinking())
+ self.assertEqual(obj.tokens_in_think, 1)
+ # Grammar should be rolled back by 3 (only generation tokens)
+ inner_grammar.rollback.assert_called_once_with(3)
+
+ def test_rollback_generation_tokens_only(self):
+ obj, inner_grammar = self._make_object_with_mock_grammar()
+ obj.maybe_init_reasoning(True)
+
+ obj.accept_token(10) # think
+ obj.accept_token(7) # think_end_id
+ obj.accept_token(20) # gen 1
+ obj.accept_token(21) # gen 2
+
+ # Rollback 1: should only roll back 1 generation token
+ obj.rollback(1)
+ self.assertTrue(obj._is_generation())
+ self.assertEqual(obj.tokens_after_end, 1)
+ inner_grammar.rollback.assert_called_once_with(1)
+
+ def test_rollback_thinking_tokens_does_not_touch_grammar(self):
+ obj, inner_grammar = self._make_object_with_mock_grammar()
+ obj.maybe_init_reasoning(True)
+
+ obj.accept_token(10)
+ obj.accept_token(11)
+ obj.accept_token(12)
+
+ obj.rollback(2)
+ self.assertTrue(obj._is_thinking())
+ self.assertEqual(obj.tokens_in_think, 1)
+ inner_grammar.rollback.assert_not_called()
+ inner_grammar.accept_token.assert_not_called()
+
+ def test_copy_preserves_state(self):
+ obj, inner_grammar = self._make_object_with_mock_grammar()
+ obj.maybe_init_reasoning(True)
+
+ obj.accept_token(10)
+ obj.accept_token(7) # think_end_id → GENERATION
+ obj.accept_token(20)
+
+ self.assertEqual(obj.tokens_in_think, 1)
+ self.assertEqual(obj.tokens_after_end, 1)
+
+ copy = obj.copy()
+ # State counters must be preserved for speculative decoding
+ self.assertEqual(copy.tokens_in_think, 1)
+ self.assertEqual(copy.tokens_after_end, 1)
+ self.assertTrue(copy._is_generation())
+ self.assertIsNotNone(copy.grammar)
+ inner_grammar.copy.assert_called_once()
+
+ def test_copy_preserves_thinking_state(self):
+ obj, inner_grammar = self._make_object_with_mock_grammar()
+ obj.maybe_init_reasoning(True)
+
+ obj.accept_token(10)
+ obj.accept_token(11)
+
+ copy = obj.copy()
+ self.assertEqual(copy.tokens_in_think, 2)
+ self.assertEqual(copy.tokens_after_end, -1)
+ self.assertTrue(copy._is_thinking())
+
+
+class TestReasonerGrammarObjectFillVocabMask(unittest.TestCase):
+ """Tests for fill_vocab_mask behavior in different states."""
+
+ def test_thinking_phase_does_not_consult_inner_grammar(self):
+ inner_grammar = MagicMock()
+ # Must return a real tensor for allocate_vocab_mask since fill_vocab_mask
+ # delegates to allocate_vocab_mask via self.grammar when grammar is not None
+ inner_grammar.allocate_vocab_mask.side_effect = lambda vs, bs, d: torch.zeros(
+ (bs, (vs + 31) // 32), dtype=torch.int32
+ )
+ obj = ReasonerGrammarObject(
+ grammar=inner_grammar,
+ think_end_id=7,
+ think_excluded_token_ids=[3, 5],
+ max_think_tokens=-1,
+ enable_token_filter=True,
+ token_filter_fn=set_token_filter_torch,
+ allocate_vocab_mask_fn=lambda vs, bs, d: torch.zeros(
+ (bs, (vs + 31) // 32), dtype=torch.int32
+ ),
+ move_vocab_mask_fn=lambda vm, d: vm,
+ apply_vocab_mask_fn=lambda l, vm: None,
+ )
+ obj.maybe_init_reasoning(True)
+ mask = obj.allocate_vocab_mask(64, 1, "cpu")
+
+ obj.fill_vocab_mask(mask, 0)
+
+ inner_grammar.fill_vocab_mask.assert_not_called()
+ # Excluded tokens (3, 5) should be blocked
+ allowed = _allowed_token_ids(mask, [0, 1, 3, 5, 7, 8])
+ self.assertEqual(allowed, [0, 1, 7, 8])
+
+ def test_generation_phase_consults_inner_grammar(self):
+ inner_grammar = MagicMock()
+ inner_grammar.allocate_vocab_mask.side_effect = lambda vs, bs, d: torch.zeros(
+ (bs, (vs + 31) // 32), dtype=torch.int32
+ )
+ obj = ReasonerGrammarObject(
+ grammar=inner_grammar,
+ think_end_id=7,
+ think_excluded_token_ids=[3, 5],
+ max_think_tokens=-1,
+ enable_token_filter=True,
+ token_filter_fn=set_token_filter_torch,
+ allocate_vocab_mask_fn=lambda vs, bs, d: torch.zeros(
+ (bs, (vs + 31) // 32), dtype=torch.int32
+ ),
+ move_vocab_mask_fn=lambda vm, d: vm,
+ apply_vocab_mask_fn=lambda l, vm: None,
+ )
+ obj.maybe_init_reasoning(True)
+ obj.accept_token(10)
+ obj.accept_token(7) # think_end_id → GENERATION
+
+ mask = obj.allocate_vocab_mask(64, 1, "cpu")
+ obj.fill_vocab_mask(mask, 0)
+
+ inner_grammar.fill_vocab_mask.assert_called_once_with(mask, 0)
+
+ def test_non_strict_thinking_is_noop(self):
+ inner_grammar = MagicMock()
+ obj = ReasonerGrammarObject(
+ grammar=inner_grammar,
+ think_end_id=7,
+ think_excluded_token_ids=None,
+ max_think_tokens=-1,
+ enable_token_filter=False,
+ token_filter_fn=None,
+ )
+ obj.maybe_init_reasoning(True)
+ mask = torch.zeros((1, 2), dtype=torch.int32)
+
+ obj.fill_vocab_mask(mask, 0)
+
+ inner_grammar.fill_vocab_mask.assert_not_called()
+ # Mask should remain all zeros (no filtering)
+ self.assertTrue(torch.all(mask == 0))
if __name__ == "__main__":
diff --git a/test/registered/unit/constrained/test_token_filter_ops.py b/test/registered/unit/constrained/test_token_filter_ops.py
new file mode 100644
index 000000000..08dba0b4b
--- /dev/null
+++ b/test/registered/unit/constrained/test_token_filter_ops.py
@@ -0,0 +1,146 @@
+"""
+Unit tests for token filter operations (Triton and Torch paths).
+
+Verifies that both implementations produce identical bitmask output
+for the same inputs, ensuring parity across GPU and CPU paths.
+"""
+
+import unittest
+
+import torch
+
+from sglang.srt.constrained.torch_ops.token_filter_torch_ops import (
+ set_token_filter_torch,
+)
+from sglang.test.ci.ci_register import register_cpu_ci
+
+register_cpu_ci(2.0, "stage-a-test-cpu")
+
+# Conditionally import Triton path
+_has_cuda = torch.cuda.is_available()
+if _has_cuda:
+ from sglang.srt.constrained.triton_ops.token_filter_ops import (
+ set_token_filter_triton,
+ )
+
+
+def _get_allowed_tokens(vocab_mask, batch_idx, max_token_id):
+ """Extract allowed token IDs from a bitmask row."""
+ allowed = []
+ for token_id in range(max_token_id):
+ elem = token_id // 32
+ bit = token_id % 32
+ val = int(vocab_mask[batch_idx, elem].item())
+ if val & (1 << bit):
+ allowed.append(token_id)
+ return allowed
+
+
+class TestSetTokenFilterTorch(unittest.TestCase):
+ """Tests for the Torch token filter implementation."""
+
+ def test_allow_tokens_from_blank_mask(self):
+ vocab_mask = torch.zeros((1, 4), dtype=torch.int32) # 128 tokens
+ set_token_filter_torch(vocab_mask, [0, 5, 31, 32, 63], 0, is_allowed=True)
+
+ allowed = _get_allowed_tokens(vocab_mask, 0, 64)
+ self.assertEqual(allowed, [0, 5, 31, 32, 63])
+
+ def test_block_tokens_from_full_mask(self):
+ vocab_mask = torch.full((1, 4), -1, dtype=torch.int32) # all bits set
+ set_token_filter_torch(
+ vocab_mask, [3, 5], 0, is_allowed=False, reset_vocab_mask=False
+ )
+
+ allowed = _get_allowed_tokens(vocab_mask, 0, 64)
+ self.assertNotIn(3, allowed)
+ self.assertNotIn(5, allowed)
+ self.assertIn(0, allowed)
+ self.assertIn(1, allowed)
+
+ def test_reset_then_allow(self):
+ vocab_mask = torch.full((1, 2), -1, dtype=torch.int32)
+ set_token_filter_torch(
+ vocab_mask, [7], 0, is_allowed=True, reset_vocab_mask=True
+ )
+
+ allowed = _get_allowed_tokens(vocab_mask, 0, 64)
+ self.assertEqual(allowed, [7])
+
+ def test_reset_then_block(self):
+ vocab_mask = torch.zeros((1, 2), dtype=torch.int32)
+ set_token_filter_torch(
+ vocab_mask, [3, 5], 0, is_allowed=False, reset_vocab_mask=True
+ )
+
+ allowed = _get_allowed_tokens(vocab_mask, 0, 64)
+ self.assertNotIn(3, allowed)
+ self.assertNotIn(5, allowed)
+ # All other tokens should be allowed (reset to -1 for block mode)
+ self.assertIn(0, allowed)
+ self.assertIn(7, allowed)
+
+ def test_empty_token_list(self):
+ vocab_mask = torch.zeros((1, 2), dtype=torch.int32)
+ set_token_filter_torch(
+ vocab_mask, [], 0, is_allowed=True, reset_vocab_mask=True
+ )
+
+ allowed = _get_allowed_tokens(vocab_mask, 0, 64)
+ self.assertEqual(allowed, [])
+
+ def test_batch_indexing(self):
+ vocab_mask = torch.zeros((3, 2), dtype=torch.int32)
+ set_token_filter_torch(vocab_mask, [1], 0, is_allowed=True)
+ set_token_filter_torch(vocab_mask, [2], 1, is_allowed=True)
+ set_token_filter_torch(vocab_mask, [3], 2, is_allowed=True)
+
+ self.assertEqual(_get_allowed_tokens(vocab_mask, 0, 64), [1])
+ self.assertEqual(_get_allowed_tokens(vocab_mask, 1, 64), [2])
+ self.assertEqual(_get_allowed_tokens(vocab_mask, 2, 64), [3])
+
+
+@unittest.skipUnless(_has_cuda, "CUDA not available")
+class TestTritonTorchParity(unittest.TestCase):
+ """Tests that Triton and Torch produce identical output."""
+
+ def _compare_outputs(self, token_ids, is_allowed, reset):
+ vocab_size = 128
+ num_elements = (vocab_size + 31) // 32
+
+ torch_mask = torch.zeros((1, num_elements), dtype=torch.int32)
+ triton_mask = torch.zeros((1, num_elements), dtype=torch.int32, device="cuda")
+
+ set_token_filter_torch(
+ torch_mask,
+ token_ids,
+ 0,
+ is_allowed=is_allowed,
+ reset_vocab_mask=reset,
+ )
+ set_token_filter_triton(
+ triton_mask,
+ token_ids,
+ 0,
+ is_allowed=is_allowed,
+ reset_vocab_mask=reset,
+ )
+
+ triton_cpu = triton_mask.cpu()
+ self.assertTrue(
+ torch.equal(torch_mask, triton_cpu),
+ f"Mismatch: torch={torch_mask} triton={triton_cpu}",
+ )
+
+ def test_parity_allow_tokens(self):
+ self._compare_outputs([0, 5, 31, 32, 63, 100], is_allowed=True, reset=True)
+
+ def test_parity_block_tokens(self):
+ self._compare_outputs([3, 5, 10], is_allowed=False, reset=True)
+
+ def test_parity_empty_tokens(self):
+ self._compare_outputs([], is_allowed=True, reset=True)
+
+
+if __name__ == "__main__":
+ unittest.main()