Skip to content

Commit 766af00

Browse files
committed
fix(sglang): pass reasoning gate for guided decoding
Bridge Dynamo's guided-decoding reasoning state into SGLang by forwarding require_reasoning from the frontend/preprocessor and setting GenerateReqInput.require_reasoning before SGLang initializes its grammar backend. Mirror Dynamo reasoning-parser names into SGLang parser names, cover the unified SglangLLMEngine path plus legacy prefill/decode handlers, and add request-gating coverage for thinking opt-out cases. Signed-off-by: Yuting Wu (DLAlgo) <yutwu@nvidia.com>
1 parent e7eb1c5 commit 766af00

12 files changed

Lines changed: 589 additions & 63 deletions

File tree

components/src/dynamo/frontend/sglang_processor.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -158,6 +158,7 @@ def _preprocess_worker(
158158
eos_token_id,
159159
pre.guided_decoding,
160160
pre.tool_call_parser,
161+
require_reasoning=pre.force_reasoning,
161162
)
162163

163164
effective_reasoning_parser_name = (
@@ -180,6 +181,7 @@ def _build_dynamo_preproc(
180181
eos_token_id: int | None,
181182
guided_decoding: dict[str, Any] | None = None,
182183
tool_call_parser: ToolCallParserType | None = None,
184+
require_reasoning: bool = False,
183185
) -> dict[str, Any]:
184186
"""Build the Dynamo preprocessed request dict from request fields."""
185187
max_tokens = request.get("max_completion_tokens") or request.get("max_tokens")
@@ -250,6 +252,11 @@ def _build_dynamo_preproc(
250252
if mm_data:
251253
preproc["multi_modal_data"] = mm_data
252254

255+
if require_reasoning:
256+
extra_args = dict(preproc.get("extra_args") or {})
257+
extra_args["require_reasoning"] = True
258+
preproc["extra_args"] = extra_args
259+
253260
return preproc
254261

255262

@@ -368,6 +375,7 @@ async def _generator_inner(
368375
self.eos_token_id,
369376
pre.guided_decoding,
370377
pre.tool_call_parser,
378+
require_reasoning=pre.force_reasoning,
371379
)
372380
except InvalidArgument:
373381
raise

components/src/dynamo/frontend/tests/test_sglang_processor_unit.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -112,6 +112,18 @@ def test_top_k_positive_preserved(self):
112112
)
113113
assert result["sampling_options"]["top_k"] == 50
114114

115+
def test_require_reasoning_forwards_backend_extra_args(self):
116+
"""Reasoning-gated guided decoding must reach the SGLang backend."""
117+
result = _build_dynamo_preproc(
118+
{"model": "test"},
119+
prompt_token_ids=[1],
120+
model_name="test",
121+
eos_token_id=None,
122+
require_reasoning=True,
123+
)
124+
assert result["extra_args"]["require_reasoning"] is True
125+
assert "prompt_injected_reasoning" not in result["extra_args"]
126+
115127
def test_sampling_options_from_request(self):
116128
"""All sampling fields are projected from request."""
117129
request = {

components/src/dynamo/sglang/args.py

Lines changed: 52 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,24 @@
3030
configure_dynamo_logging()
3131

3232

33+
_DYN_TO_SGLANG_REASONING_PARSER = {
34+
"deepseek_r1": "deepseek-r1",
35+
"deepseek_v3": "deepseek-v3",
36+
"deepseek_v3_1": "deepseek-v3",
37+
"deepseek_v3_2": "deepseek-v3",
38+
"deepseek_v4": "deepseek-v4",
39+
"deepseekv4": "deepseek-v4",
40+
"gemma-4": "gemma4",
41+
"gpt_oss": "gpt-oss",
42+
"kimi_k25": "kimi_k2",
43+
"minimax_append_think": "minimax-append-think",
44+
"nemotron_nano": "nemotron_3",
45+
"nemotron3": "nemotron_3",
46+
"nemotron_v3": "nemotron_3",
47+
"nemotron_deci": "glm45",
48+
}
49+
50+
3351
class DynamoConfig(DynamoRuntimeConfig, DynamoSGLangConfig):
3452
"""Combined configuration container for SGLang server and Dynamo args."""
3553

@@ -79,12 +97,44 @@ def _preprocess_for_encode_config(
7997
def _validate_parser_flags(
8098
sglang_val: Optional[str], dynamo_val: Optional[str], name: str
8199
) -> None:
82-
"""Validate that --{name} (SGLang) and --dyn-{name} (Dynamo) are not both set."""
100+
"""Validate parser flag combinations."""
101+
if name == "reasoning-parser" and sglang_val and dynamo_val:
102+
expected_sglang_val = _DYN_TO_SGLANG_REASONING_PARSER.get(
103+
dynamo_val, dynamo_val
104+
)
105+
if sglang_val != expected_sglang_val:
106+
logging.error(
107+
"Cannot use different --reasoning-parser (%s) and "
108+
"--dyn-reasoning-parser (%s) values. Expected SGLang parser "
109+
"%s for Dynamo parser %s.",
110+
sglang_val,
111+
dynamo_val,
112+
expected_sglang_val,
113+
dynamo_val,
114+
)
115+
sys.exit(1)
116+
return
83117
if sglang_val and dynamo_val:
84118
logging.error(f"Cannot use both --{name} and --dyn-{name}.")
85119
sys.exit(1)
86120

87121

122+
def _mirror_dyn_reasoning_parser_to_sglang(
123+
parsed_args: argparse.Namespace, dynamo_config: DynamoSGLangConfig
124+
) -> None:
125+
"""Use Dynamo's reasoning parser to activate SGLang's reasoner gate."""
126+
dyn_parser = dynamo_config.dyn_reasoning_parser
127+
if dyn_parser and not getattr(parsed_args, "reasoning_parser", None):
128+
sglang_parser = _DYN_TO_SGLANG_REASONING_PARSER.get(dyn_parser, dyn_parser)
129+
parsed_args.reasoning_parser = sglang_parser
130+
logging.info(
131+
"Mirroring --dyn-reasoning-parser=%s to SGLang --reasoning-parser=%s "
132+
"so guided decoding is gated until the reasoning end token.",
133+
dyn_parser,
134+
sglang_parser,
135+
)
136+
137+
88138
def _has_cli_flag(args: list[str], flag: str) -> bool:
89139
"""Return True when a CLI flag is present in '--flag val' or '--flag=val' form."""
90140
return any(arg == flag or arg.startswith(f"{flag}=") for arg in args)
@@ -295,6 +345,7 @@ async def parse_args(args: list[str]) -> Config:
295345
dynamo_config.dyn_reasoning_parser,
296346
"reasoning-parser",
297347
)
348+
_mirror_dyn_reasoning_parser_to_sglang(parsed_args, dynamo_config)
298349

299350
if dynamo_config.custom_jinja_template and dynamo_config.use_sglang_tokenizer:
300351
logging.error(

components/src/dynamo/sglang/llm_engine.py

Lines changed: 27 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -59,6 +59,10 @@
5959
runtime_capacity,
6060
)
6161
from dynamo.sglang.publisher import format_zmq_endpoint
62+
from dynamo.sglang.reasoning import (
63+
install_require_reasoning_proxy,
64+
require_reasoning_context,
65+
)
6266

6367
if TYPE_CHECKING:
6468
from dynamo._core.backend import EngineMetrics # type: ignore[import-not-found]
@@ -149,6 +153,7 @@ async def start(self, worker_id: int) -> EngineConfig:
149153
del worker_id # SGLang bootstrap uses host/port/room triples
150154

151155
self.engine = sgl.Engine(server_args=self.server_args)
156+
install_require_reasoning_proxy(self.engine)
152157

153158
tokenizer = (
154159
self.engine.tokenizer_manager.tokenizer
@@ -289,19 +294,22 @@ async def generate(
289294
"SGLang",
290295
)
291296

292-
stream = await self.engine.async_generate(
293-
**input_param,
294-
sampling_params=sampling_params,
295-
stream=True,
296-
rid=context.trace_id,
297-
data_parallel_rank=sgl_dp_rank,
298-
**telemetry.engine_trace_kwargs(
299-
context,
300-
kwarg_name="external_trace_header",
301-
enabled=self.enable_trace,
302-
),
303-
**bootstrap_kwargs,
304-
)
297+
with require_reasoning_context(
298+
self._has_reasoning_parser(), dict(request), input_param
299+
):
300+
stream = await self.engine.async_generate(
301+
**input_param,
302+
sampling_params=sampling_params,
303+
stream=True,
304+
rid=context.trace_id,
305+
data_parallel_rank=sgl_dp_rank,
306+
**telemetry.engine_trace_kwargs(
307+
context,
308+
kwarg_name="external_trace_header",
309+
enabled=self.enable_trace,
310+
),
311+
**bootstrap_kwargs,
312+
)
305313

306314
# ORDER MATTERS: async_generate must register the room (the await
307315
# above) before we yield the bootstrap chunk — otherwise the
@@ -719,3 +727,9 @@ def _get_input_param(self, request: GenerateRequest) -> dict:
719727
return {
720728
"prompt" if isinstance(request_input, str) else "input_ids": request_input
721729
}
730+
731+
def _has_reasoning_parser(self) -> bool:
732+
return bool(
733+
getattr(self.server_args, "reasoning_parser", None)
734+
or getattr(self.dynamo_args, "dyn_reasoning_parser", None)
735+
)
Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,65 @@
1+
# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2+
# SPDX-License-Identifier: Apache-2.0
3+
4+
import contextvars
5+
import functools
6+
import threading
7+
from contextlib import contextmanager
8+
from typing import Any, Dict
9+
10+
_DYN_REQUIRE_REASONING_CV: contextvars.ContextVar[bool] = contextvars.ContextVar(
11+
"dynamo_sglang_require_reasoning", default=False
12+
)
13+
_REQUIRE_REASONING_PROXY_LOCK = threading.Lock()
14+
15+
16+
def install_require_reasoning_proxy(engine: Any) -> None:
17+
"""Set SGLang GenerateReqInput.require_reasoning from a per-request context."""
18+
tm = getattr(engine, "tokenizer_manager", None)
19+
if tm is None or getattr(tm, "_dynamo_require_reasoning_wrapped", False):
20+
return
21+
22+
with _REQUIRE_REASONING_PROXY_LOCK:
23+
if getattr(tm, "_dynamo_require_reasoning_wrapped", False):
24+
return
25+
26+
original = tm.generate_request
27+
28+
@functools.wraps(original)
29+
def _wrapped(obj, request):
30+
if _DYN_REQUIRE_REASONING_CV.get():
31+
obj.require_reasoning = True
32+
return original(obj, request)
33+
34+
tm._dynamo_require_reasoning_wrapped = True # type: ignore[attr-defined]
35+
tm.generate_request = _wrapped # type: ignore[assignment]
36+
37+
38+
def request_requires_reasoning(
39+
has_reasoning_parser: bool, request: Dict[str, Any], input_param: Dict[str, Any]
40+
) -> bool:
41+
if not has_reasoning_parser:
42+
return False
43+
44+
extra_args = request.get("extra_args") or {}
45+
if isinstance(extra_args, dict) and (
46+
extra_args.get("require_reasoning")
47+
or extra_args.get("prompt_injected_reasoning")
48+
):
49+
return True
50+
51+
prompt = input_param.get("prompt")
52+
return isinstance(prompt, str) and prompt.rstrip().endswith("<think>")
53+
54+
55+
@contextmanager
56+
def require_reasoning_context(
57+
has_reasoning_parser: bool, request: Dict[str, Any], input_param: Dict[str, Any]
58+
):
59+
token = _DYN_REQUIRE_REASONING_CV.set(
60+
request_requires_reasoning(has_reasoning_parser, request, input_param)
61+
)
62+
try:
63+
yield
64+
finally:
65+
_DYN_REQUIRE_REASONING_CV.reset(token)

components/src/dynamo/sglang/request_handlers/handler_base.py

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,10 @@
4444
from dynamo.runtime import DistributedRuntime
4545
from dynamo.sglang.args import Config
4646
from dynamo.sglang.publisher import DynamoSglangPublisher
47+
from dynamo.sglang.reasoning import (
48+
install_require_reasoning_proxy,
49+
require_reasoning_context,
50+
)
4751

4852
logger = logging.getLogger(__name__)
4953

@@ -720,6 +724,7 @@ def __init__(
720724
self.enable_trace = getattr(config.server_args, "enable_trace", False)
721725

722726
if engine is not None:
727+
install_require_reasoning_proxy(engine)
723728
self.input_param_manager = InputParamManager(
724729
self.engine.tokenizer_manager.tokenizer
725730
if self.use_sglang_tokenizer
@@ -1059,6 +1064,19 @@ def _get_input_param(self, request: Dict[str, Any]) -> Dict[str, Any]:
10591064
"prompt" if isinstance(request_input, str) else "input_ids": request_input
10601065
}
10611066

1067+
def _has_reasoning_parser(self) -> bool:
1068+
return bool(
1069+
getattr(self.config.server_args, "reasoning_parser", None)
1070+
or getattr(self.config.dynamo_args, "dyn_reasoning_parser", None)
1071+
)
1072+
1073+
def _require_reasoning_context(
1074+
self, request: Dict[str, Any], input_param: Dict[str, Any]
1075+
):
1076+
return require_reasoning_context(
1077+
self._has_reasoning_parser(), request, input_param
1078+
)
1079+
10621080
def _session_kwargs(self, request: Dict[str, Any]) -> Dict[str, Any]:
10631081
if not getattr(self.config.server_args, "enable_streaming_session", False):
10641082
return {}

components/src/dynamo/sglang/request_handlers/llm/decode_handler.py

Lines changed: 33 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -466,22 +466,23 @@ async def generate(
466466
routing = request.get("routing") or {}
467467
dp_rank = routing.get("dp_rank")
468468

469-
decode = await self.engine.async_generate(
470-
**input_param,
471-
sampling_params=sampling_params,
472-
stream=True,
473-
**self._routed_experts_kwargs,
474-
bootstrap_host=bootstrap_info["bootstrap_host"],
475-
bootstrap_port=bootstrap_info["bootstrap_port"],
476-
bootstrap_room=bootstrap_info["bootstrap_room"],
477-
external_trace_header=trace_header,
478-
rid=trace_id,
479-
data_parallel_rank=dp_rank,
480-
**self._session_kwargs(request),
481-
lora_path=lora_path,
482-
**logprob_kwargs,
483-
**self._priority_kwargs(priority),
484-
)
469+
with self._require_reasoning_context(request, input_param):
470+
decode = await self.engine.async_generate(
471+
**input_param,
472+
sampling_params=sampling_params,
473+
stream=True,
474+
**self._routed_experts_kwargs,
475+
bootstrap_host=bootstrap_info["bootstrap_host"],
476+
bootstrap_port=bootstrap_info["bootstrap_port"],
477+
bootstrap_room=bootstrap_info["bootstrap_room"],
478+
external_trace_header=trace_header,
479+
rid=trace_id,
480+
data_parallel_rank=dp_rank,
481+
**self._session_kwargs(request),
482+
lora_path=lora_path,
483+
**logprob_kwargs,
484+
**self._priority_kwargs(priority),
485+
)
485486

486487
if not self.use_sglang_tokenizer:
487488
async for out in self._process_token_stream(
@@ -525,21 +526,22 @@ async def generate(
525526
routing = request.get("routing") or {}
526527
dp_rank = routing.get("dp_rank")
527528

528-
agg = await self.engine.async_generate(
529-
**input_param,
530-
image_data=image_data,
531-
video_data=video_data,
532-
sampling_params=sampling_params,
533-
stream=True,
534-
**self._routed_experts_kwargs,
535-
external_trace_header=trace_header,
536-
rid=trace_id,
537-
data_parallel_rank=dp_rank,
538-
**self._session_kwargs(request),
539-
lora_path=lora_path,
540-
**logprob_kwargs,
541-
**self._priority_kwargs(priority),
542-
)
529+
with self._require_reasoning_context(request, input_param):
530+
agg = await self.engine.async_generate(
531+
**input_param,
532+
image_data=image_data,
533+
video_data=video_data,
534+
sampling_params=sampling_params,
535+
stream=True,
536+
**self._routed_experts_kwargs,
537+
external_trace_header=trace_header,
538+
rid=trace_id,
539+
data_parallel_rank=dp_rank,
540+
**self._session_kwargs(request),
541+
lora_path=lora_path,
542+
**logprob_kwargs,
543+
**self._priority_kwargs(priority),
544+
)
543545
if not self.use_sglang_tokenizer:
544546
async for out in self._process_token_stream(
545547
agg,

0 commit comments

Comments
 (0)