From cbca049fd49d7e615ff8af3bcf6c1dcdea790745 Mon Sep 17 00:00:00 2001 From: Jan Hilgard Date: Mon, 6 Apr 2026 23:32:38 +0200 Subject: [PATCH] Add per-request enable_thinking API parameter Allows controlling thinking/reasoning mode per-request via the enable_thinking field in extra_body. Three-level priority: 1. Per-request: extra_body.enable_thinking (true/false) 2. Environment: VLLM_MLX_ENABLE_THINKING (true/false/1/0/yes/no) 3. Default: true All code paths (SimpleEngine, BatchedEngine, MLLM) now consistently use the VLLM_MLX_ENABLE_THINKING env var as fallback, replacing the previous "coder" model name heuristic. Co-Authored-By: Claude Opus 4.6 --- vllm_mlx/api/models.py | 2 ++ vllm_mlx/engine/batched.py | 21 +++++++++++++++++++-- vllm_mlx/engine/simple.py | 18 ++++++++++++------ vllm_mlx/models/mllm.py | 12 ++++++++++++ vllm_mlx/server.py | 11 ++++++++--- 5 files changed, 53 insertions(+), 11 deletions(-) diff --git a/vllm_mlx/api/models.py b/vllm_mlx/api/models.py index 32b26e035..38ca91639 100644 --- a/vllm_mlx/api/models.py +++ b/vllm_mlx/api/models.py @@ -178,6 +178,8 @@ class ChatCompletionRequest(BaseModel): specprefill: bool | None = None # SpecPrefill: per-request keep percentage (0.0-1.0, None = use server default) specprefill_keep_pct: float | None = None + # Enable/disable thinking mode (None = server default, typically True) + enable_thinking: bool | None = None class AssistantMessage(BaseModel): diff --git a/vllm_mlx/engine/batched.py b/vllm_mlx/engine/batched.py index 3ac52b4b0..d1c265d01 100644 --- a/vllm_mlx/engine/batched.py +++ b/vllm_mlx/engine/batched.py @@ -12,6 +12,7 @@ """ import logging +import os from collections.abc import AsyncIterator from typing import Any @@ -335,6 +336,7 @@ def _apply_chat_template( messages: list[dict[str, Any]], tools: list[dict] | None = None, num_images: int = 0, + enable_thinking: bool | None = None, ) -> str: """Apply chat template to messages. @@ -363,9 +365,15 @@ def _apply_chat_template( if self._is_mllm and num_images > 0: messages = self._prepare_mllm_messages(messages) + # Per-request enable_thinking override; fall back to env var / default True. + if enable_thinking is None: + enable_thinking = os.environ.get( + "VLLM_MLX_ENABLE_THINKING", "true" + ).lower() in ("true", "1", "yes") template_kwargs = { "tokenize": False, "add_generation_prompt": True, + "enable_thinking": enable_thinking, } if tools: template_kwargs["tools"] = tools @@ -375,9 +383,10 @@ def _apply_chat_template( messages, **template_kwargs ) except TypeError as e: - # Some templates don't accept 'tools'; retry without them. + # Some templates don't accept 'tools' or 'enable_thinking'; + # retry without them. logger.debug(f"Chat template TypeError, retrying without extras: {e}") - for key in ["tools"]: + for key in ["tools", "enable_thinking"]: if key in template_kwargs: del template_kwargs[key] return template_applicator.apply_chat_template( @@ -621,11 +630,15 @@ async def chat( # Convert tools for template template_tools = convert_tools_for_template(tools) if tools else None + # Per-request enable_thinking override + enable_thinking = kwargs.pop("enable_thinking", None) + # Apply chat template prompt = self._apply_chat_template( messages, template_tools, num_images=len(all_images), + enable_thinking=enable_thinking, ) return await self.generate( @@ -732,11 +745,15 @@ async def stream_chat( # Convert tools for template template_tools = convert_tools_for_template(tools) if tools else None + # Per-request enable_thinking override + enable_thinking = kwargs.pop("enable_thinking", None) + # Apply chat template prompt = self._apply_chat_template( messages, template_tools, num_images=len(all_images), + enable_thinking=enable_thinking, ) # Compute prefix boundary for cache diff --git a/vllm_mlx/engine/simple.py b/vllm_mlx/engine/simple.py index da3ccfc18..6146171cb 100644 --- a/vllm_mlx/engine/simple.py +++ b/vllm_mlx/engine/simple.py @@ -8,6 +8,7 @@ import asyncio import logging +import os from collections.abc import AsyncIterator from typing import Any @@ -589,9 +590,12 @@ def run_stream(): # For LLM, apply chat template and stream tokenizer = self._model.tokenizer if hasattr(tokenizer, "apply_chat_template"): - # Disable thinking mode for coder models since it interferes - # with tool call parsing (tags leak as raw text). - enable_thinking = "coder" not in self._model_name.lower() + # Per-request enable_thinking override; fall back to env var / default True. + enable_thinking = kwargs.pop("enable_thinking", None) + if enable_thinking is None: + enable_thinking = os.environ.get( + "VLLM_MLX_ENABLE_THINKING", "true" + ).lower() in ("true", "1", "yes") template_kwargs = { "tokenize": False, "add_generation_prompt": True, @@ -835,9 +839,11 @@ async def _stream_generate_text( specprefill_override = kwargs.pop("specprefill", None) specprefill_keep_pct = kwargs.pop("specprefill_keep_pct", None) - # Read enable_thinking from env (set by runtime_patches, consistent with MLLM path) - enable_thinking_env = os.environ.get("VLLM_MLX_ENABLE_THINKING", "true") - enable_thinking = enable_thinking_env.lower() in ("true", "1", "yes") + # Per-request enable_thinking override; fall back to env var / default True. + enable_thinking = kwargs.pop("enable_thinking", None) + if enable_thinking is None: + enable_thinking_env = os.environ.get("VLLM_MLX_ENABLE_THINKING", "true") + enable_thinking = enable_thinking_env.lower() in ("true", "1", "yes") # Apply chat template for full prompt template_kwargs = { diff --git a/vllm_mlx/models/mllm.py b/vllm_mlx/models/mllm.py index fcf3537f4..50ff5e473 100644 --- a/vllm_mlx/models/mllm.py +++ b/vllm_mlx/models/mllm.py @@ -1328,6 +1328,11 @@ def chat( video_max_frames = kwargs.pop("video_max_frames", MAX_FRAMES) tools = kwargs.pop("tools", None) use_cache = kwargs.pop("use_cache", True) + enable_thinking = kwargs.pop("enable_thinking", None) + if enable_thinking is None: + enable_thinking = os.environ.get( + "VLLM_MLX_ENABLE_THINKING", "true" + ).lower() in ("true", "1", "yes") # Collect video inputs from messages _msg_video_inputs = self._collect_video_inputs(messages) @@ -1458,6 +1463,7 @@ def chat( self.processor, chat_messages, add_generation_prompt=True, + enable_thinking=enable_thinking, **template_extra_kwargs, ) except Exception as e: @@ -1724,6 +1730,11 @@ def stream_chat( video_max_frames = kwargs.pop("video_max_frames", MAX_FRAMES) tools = kwargs.pop("tools", None) use_cache = kwargs.pop("use_cache", True) + enable_thinking = kwargs.pop("enable_thinking", None) + if enable_thinking is None: + enable_thinking = os.environ.get( + "VLLM_MLX_ENABLE_THINKING", "true" + ).lower() in ("true", "1", "yes") # Collect video inputs from messages _msg_video_inputs = self._collect_video_inputs(messages) @@ -1838,6 +1849,7 @@ def stream_chat( self.processor, chat_messages, add_generation_prompt=True, + enable_thinking=enable_thinking, **template_extra_kwargs, ) except Exception as e: diff --git a/vllm_mlx/server.py b/vllm_mlx/server.py index af10e7341..184b032db 100644 --- a/vllm_mlx/server.py +++ b/vllm_mlx/server.py @@ -1423,6 +1423,10 @@ async def create_chat_completion(request: ChatCompletionRequest, raw_request: Re if request.specprefill_keep_pct is not None: chat_kwargs["specprefill_keep_pct"] = request.specprefill_keep_pct + # Enable/disable thinking mode per request + if request.enable_thinking is not None: + chat_kwargs["enable_thinking"] = request.enable_thinking + # Add tools if provided if request.tools: chat_kwargs["tools"] = convert_tools_for_template(request.tools) @@ -1458,8 +1462,9 @@ async def create_chat_completion(request: ChatCompletionRequest, raw_request: Re cleaned_text, tool_calls = _parse_tool_calls_with_parser(output.text, request) # Extract reasoning content FIRST (strips channel tokens before JSON extraction) + # Skip reasoning parser when enable_thinking=False (no think tags expected) reasoning_text = None - if _reasoning_parser and not tool_calls: + if _reasoning_parser and not tool_calls and request.enable_thinking is not False: text_to_parse = cleaned_text or output.text reasoning_text, cleaned_text = _reasoning_parser.extract_reasoning( text_to_parse @@ -2082,8 +2087,8 @@ async def stream_chat_completion( if hasattr(output, "completion_tokens") and output.completion_tokens: completion_tokens = output.completion_tokens - # Use reasoning parser if enabled - if _reasoning_parser and delta_text: + # Use reasoning parser if enabled (skip when enable_thinking=False) + if _reasoning_parser and delta_text and request.enable_thinking is not False: previous_text = accumulated_text accumulated_text += delta_text delta_msg = _reasoning_parser.extract_reasoning_streaming(