feat(inference): add use_adapter field for per-request adapter toggling in compare mode
This commit is contained in:
parent
6ed179b459
commit
f67ee58347
3 changed files with 152 additions and 51 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ────────────────────────────────────
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue