From 909955767b1fe051202f3a64d13d92cb4e160d28 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Mon, 16 Feb 2026 06:33:17 +0000 Subject: [PATCH] feat: add min_p sampling parameter to /chat/completions generation pipeline --- studio/backend/core/inference/inference.py | 12 +++++++++--- studio/backend/models/inference.py | 1 + studio/backend/routes/inference.py | 1 + 3 files changed, 11 insertions(+), 3 deletions(-) diff --git a/studio/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py index 6be7e63306..38394f6f33 100644 --- a/studio/backend/core/inference/inference.py +++ b/studio/backend/core/inference/inference.py @@ -515,6 +515,7 @@ class InferenceBackend: temperature: float = 0.7, top_p: float = 0.9, top_k: int = 40, + min_p: float = 0.0, max_new_tokens: int = 256, repetition_penalty: float = 1.1, cancel_event=None) -> Generator[str, None, None]: @@ -531,6 +532,7 @@ class InferenceBackend: temperature=temperature, top_p=top_p, top_k=top_k, + min_p=min_p, max_new_tokens=max_new_tokens, repetition_penalty=repetition_penalty, cancel_event=cancel_event, @@ -543,6 +545,7 @@ class InferenceBackend: temperature: float = 0.7, top_p: float = 0.9, top_k: int = 40, + min_p: float = 0.0, max_new_tokens: int = 256, repetition_penalty: float = 1.1, cancel_event=None) -> Generator[str, None, None]: @@ -562,7 +565,7 @@ class InferenceBackend: # Vision model generation yield from self._generate_vision_response( messages, system_prompt, image, - temperature, top_p, top_k, max_new_tokens, repetition_penalty, + temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty, cancel_event=cancel_event, ) else: @@ -605,12 +608,12 @@ class InferenceBackend: # Step 3: Generate yield from self.generate_stream( - formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty, + formatted_prompt, temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty, cancel_event=cancel_event, ) def _generate_vision_response(self, messages, system_prompt, image, - temperature, top_p, top_k, max_new_tokens, + temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty, cancel_event=None) -> Generator[str, None, None]: """Handle vision model generation with true token-by-token streaming.""" model_info = self.models[self.active_model_name] @@ -671,6 +674,7 @@ class InferenceBackend: temperature=temperature, top_p=top_p, top_k=top_k, + min_p=min_p, ) err: dict[str, str] = {} @@ -728,6 +732,7 @@ class InferenceBackend: temperature: float = 0.7, top_p: float = 0.9, top_k: int = 40, + min_p: float = 0.0, max_new_tokens: int = 256, repetition_penalty: float = 1.1, cancel_event=None) -> Generator[str, None, None]: @@ -760,6 +765,7 @@ class InferenceBackend: temperature=temperature, top_p=top_p, top_k=top_k, + min_p=min_p, repetition_penalty=repetition_penalty, do_sample=True, eos_token_id=tokenizer.eos_token_id, diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index b3924c5569..1b661047ae 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -129,6 +129,7 @@ class ChatCompletionRequest(BaseModel): # ── Unsloth extensions (ignored by standard OpenAI clients) ── top_k: int = Field(40, ge=1, le=100, description="[x-unsloth] Top-k sampling") + min_p: float = Field(0.0, ge=0.0, le=1.0, description="[x-unsloth] Min-p sampling threshold") repetition_penalty: float = Field(1.1, ge=1.0, le=2.0, description="[x-unsloth] Repetition penalty") image_base64: Optional[str] = Field(None, description="[x-unsloth] Base64-encoded image for vision models") use_adapter: Optional[Union[bool, str]] = Field( diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 3bb84d5c6b..c30a1638d4 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -371,6 +371,7 @@ async def openai_chat_completions(payload: ChatCompletionRequest, request: Reque temperature=payload.temperature, top_p=payload.top_p, top_k=payload.top_k, + min_p=payload.min_p, max_new_tokens=payload.max_tokens or 512, repetition_penalty=payload.repetition_penalty, )