From f67ee5834779a7bc152ec88c3bf364bd4c29d83a Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Sat, 14 Feb 2026 14:52:13 +0000 Subject: [PATCH] feat(inference): add use_adapter field for per-request adapter toggling in compare mode --- studio/backend/core/inference/inference.py | 173 +++++++++++++++------ studio/backend/models/inference.py | 10 ++ studio/backend/routes/inference.py | 16 +- 3 files changed, 150 insertions(+), 49 deletions(-) diff --git a/studio/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py index 0fd1905ff4..f8c0213b9b 100644 --- a/studio/backend/core/inference/inference.py +++ b/studio/backend/core/inference/inference.py @@ -8,7 +8,7 @@ from peft import PeftModel, PeftModelForCausalLM import sys import torch -from typing import Optional, Generator, Tuple +from typing import Optional, Union, Generator, Tuple from utils.models import ModelConfig, get_base_model_from_lora from utils.paths import is_model_cached from utils.utils import format_error_message @@ -448,6 +448,64 @@ class InferenceBackend: return False pass + def _apply_adapter_state(self, use_adapter: Optional[Union[bool, str]]) -> None: + """ + Apply adapter state before generation. Must be called under _generation_lock. + + Args: + use_adapter: None = no change, False = disable (base model), + True = enable current adapter, str = enable specific adapter. + """ + if use_adapter is None: + return + + base = self.active_model_name + if not base or base not in self.models: + return + + model_info = self.models[base] + + if use_adapter is False: + # Disable all adapters → pure base model generation + logger.info(f"Compare mode: disabling adapters on '{base}' (base model generation)") + self.disable_adapters(base) + + elif use_adapter is True: + # Enable the most recently loaded adapter + loaded = model_info.get("loaded_adapters", {}) + if loaded: + adapter_name = list(loaded.keys())[-1] + logger.info(f"Compare mode: enabling adapter '{adapter_name}' on '{base}'") + self.set_active_adapter(base, adapter_name) + else: + logger.warning("use_adapter=true but no adapters are loaded on the model") + + elif isinstance(use_adapter, str): + # Enable a specific named adapter + logger.info(f"Compare mode: enabling specific adapter '{use_adapter}' on '{base}'") + self.set_active_adapter(base, use_adapter) + + def generate_with_adapter_control( + self, + use_adapter: Optional[Union[bool, str]] = None, + **gen_kwargs, + ) -> Generator[str, None, None]: + """ + Thread-safe generation with optional adapter toggling. + + Acquires the generation lock, applies adapter state, then generates. + This ensures adapter toggle + generation are atomic — critical for + compare mode where base and LoRA panes fire concurrently. + + Args: + use_adapter: Adapter control (None/False/True/str). See _apply_adapter_state. + **gen_kwargs: Forwarded to generate_chat_response. + """ + with self._generation_lock: + self._apply_adapter_state(use_adapter) + # Delegate to the lock-free generation path + yield from self._generate_chat_response_inner(**gen_kwargs) + def generate_chat_response(self, messages: list, system_prompt: str, @@ -459,11 +517,33 @@ class InferenceBackend: repetition_penalty: float = 1.1) -> Generator[str, None, None]: """ Generate response for text or vision models. + Acquires the generation lock. For adapter-controlled generation, + use generate_with_adapter_control() instead. + """ + with self._generation_lock: + yield from self._generate_chat_response_inner( + messages=messages, + system_prompt=system_prompt, + image=image, + temperature=temperature, + top_p=top_p, + top_k=top_k, + max_new_tokens=max_new_tokens, + repetition_penalty=repetition_penalty, + ) - 1. Messages are already in ChatML format (role/content) - 2. Apply get_chat_template() if model in mapper - 3. Apply tokenizer.apply_chat_template() - 4. Generate + def _generate_chat_response_inner(self, + messages: list, + system_prompt: str = "", + image=None, + temperature: float = 0.7, + top_p: float = 0.9, + top_k: int = 40, + max_new_tokens: int = 256, + repetition_penalty: float = 1.1) -> Generator[str, None, None]: + """ + Inner generation logic (no lock). Called by both generate_chat_response + and generate_with_adapter_control. """ if not self.active_model_name: yield "Error: No active model" @@ -473,55 +553,54 @@ class InferenceBackend: is_vision = model_info.get("is_vision", False) tokenizer = model_info.get("tokenizer") or model_info.get("processor") - with self._generation_lock: - if is_vision: - # Vision model generation - yield from self._generate_vision_response( - messages, system_prompt, image, - temperature, top_p, top_k, max_new_tokens, repetition_penalty - ) - else: - # Text model: Use training pipeline approach - # Messages are already in ChatML format from eval.py + if is_vision: + # Vision model generation + yield from self._generate_vision_response( + messages, system_prompt, image, + temperature, top_p, top_k, max_new_tokens, repetition_penalty + ) + else: + # Text model: Use training pipeline approach + # Messages are already in ChatML format from eval.py - # Step 1: Apply get_chat_template if model is in mapper - try: - from utils.datasets import MODEL_TO_TEMPLATE_MAPPER, get_tokenizer_chat_template + # Step 1: Apply get_chat_template if model is in mapper + try: + from utils.datasets import MODEL_TO_TEMPLATE_MAPPER, get_tokenizer_chat_template - model_name_lower = self.active_model_name.lower() + model_name_lower = self.active_model_name.lower() - # Check if model has a registered template - if model_name_lower in MODEL_TO_TEMPLATE_MAPPER: - template_name = MODEL_TO_TEMPLATE_MAPPER[model_name_lower] - logger.info(f"Applying chat template '{template_name}' for {self.active_model_name}") + # Check if model has a registered template + if model_name_lower in MODEL_TO_TEMPLATE_MAPPER: + template_name = MODEL_TO_TEMPLATE_MAPPER[model_name_lower] + logger.info(f"Applying chat template '{template_name}' for {self.active_model_name}") - # This modifies the tokenizer with the correct template - tokenizer = get_chat_template( - tokenizer, - self.active_model_name - ) - else: - logger.info(f"No registered template for {self.active_model_name}, using tokenizer default") - except Exception as e: - logger.warning(f"Could not apply get_chat_template: {e}") - - # Step 2: Format with tokenizer.apply_chat_template() - try: - formatted_prompt = tokenizer.apply_chat_template( - messages, - tokenize=False, - add_generation_prompt=True + # This modifies the tokenizer with the correct template + tokenizer = get_chat_template( + tokenizer, + self.active_model_name ) - logger.debug(f"Formatted prompt: {formatted_prompt[:200]}...") - except Exception as e: - logger.error(f"Error applying chat template: {e}") - # Fallback to manual formatting - formatted_prompt = self.format_chat_prompt(messages, system_prompt) + else: + logger.info(f"No registered template for {self.active_model_name}, using tokenizer default") + except Exception as e: + logger.warning(f"Could not apply get_chat_template: {e}") - # Step 3: Generate - yield from self.generate_stream( - formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty + # Step 2: Format with tokenizer.apply_chat_template() + try: + formatted_prompt = tokenizer.apply_chat_template( + messages, + tokenize=False, + add_generation_prompt=True ) + logger.debug(f"Formatted prompt: {formatted_prompt[:200]}...") + except Exception as e: + logger.error(f"Error applying chat template: {e}") + # Fallback to manual formatting + formatted_prompt = self.format_chat_prompt(messages, system_prompt) + + # Step 3: Generate + yield from self.generate_stream( + formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty + ) def _generate_vision_response(self, messages, system_prompt, image, temperature, top_p, top_k, max_new_tokens, diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index 64791b06f9..ada8bdd539 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -130,6 +130,16 @@ class ChatCompletionRequest(BaseModel): top_k: int = Field(40, ge=1, le=100, description="[x-unsloth] Top-k sampling") 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( + None, + description=( + "[x-unsloth] Adapter control for compare mode. " + "null = no change (default), " + "false = disable adapters (base model), " + "true = enable the current adapter, " + "string = enable a specific adapter by name." + ), + ) # ── Streaming response chunks ──────────────────────────────────── diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 74cf37138f..9e57946565 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -365,6 +365,18 @@ async def openai_chat_completions(request: ChatCompletionRequest): repetition_penalty=request.repetition_penalty, ) + # ── Choose generation path (adapter-controlled or standard) ── + if request.use_adapter is not None: + # Compare mode: toggle adapter state atomically with generation + def generate(): + return backend.generate_with_adapter_control( + use_adapter=request.use_adapter, **gen_kwargs + ) + else: + # Standard path: no adapter toggling + def generate(): + return backend.generate_chat_response(**gen_kwargs) + model_name = backend.active_model_name or request.model completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}" created = int(time.time()) @@ -388,7 +400,7 @@ async def openai_chat_completions(request: ChatCompletionRequest): # Content chunks — generate_chat_response yields cumulative # text, so we diff to get incremental deltas. prev_text = "" - for cumulative in backend.generate_chat_response(**gen_kwargs): + for cumulative in generate(): new_text = cumulative[len(prev_text):] prev_text = cumulative if not new_text: @@ -439,7 +451,7 @@ async def openai_chat_completions(request: ChatCompletionRequest): else: try: full_text = "" - for token in backend.generate_chat_response(**gen_kwargs): + for token in generate(): full_text = token # generate_stream yields cumulative text response = ChatCompletion(