feat(inference): add use_adapter field for per-request adapter toggling in compare mode

This commit is contained in:
Roland Tannous 2026-02-14 14:52:13 +00:00
commit f67ee58347
3 changed files with 152 additions and 51 deletions

View file

@ -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,

View file

@ -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 ────────────────────────────────────

View file

@ -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(