Merge pull request #79 from unslothai/feat/compare-use-adapter
Adapter Toggling for chat compare feature
This commit is contained in:
commit
1d362a36c3
3 changed files with 182 additions and 80 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
|
||||
|
|
@ -179,24 +179,21 @@ class InferenceBackend:
|
|||
model = self.models[base_model_name].get("model")
|
||||
|
||||
try:
|
||||
# Step 1: Unload the adapter weights. This returns the base model object.
|
||||
# This step is only necessary if the model is currently a PeftModel instance.
|
||||
# Step 1: Unload the adapter weights if model is a PeftModel.
|
||||
if isinstance(model, (PeftModel, PeftModelForCausalLM)):
|
||||
logger.info("Model is a PeftModel. Unloading adapters...")
|
||||
logger.info(f"Unloading LoRA adapters from '{base_model_name}'...")
|
||||
unwrapped_base_model = model.unload()
|
||||
self.models[base_model_name]["model"] = unwrapped_base_model
|
||||
model = unwrapped_base_model # Continue with the unwrapped model
|
||||
model = unwrapped_base_model
|
||||
|
||||
# Step 2: Delete any lingering adapter configurations from the object.
|
||||
# This is the crucial step you identified.
|
||||
if hasattr(model, 'peft_config') and model.peft_config:
|
||||
logger.info("Found lingering adapter configurations. Deleting them now...")
|
||||
# Create a static list of keys before iterating and deleting
|
||||
for name in list(model.peft_config.keys()):
|
||||
logger.info(f"Deleting adapter config: '{name}'")
|
||||
model.delete_adapter(name)
|
||||
# Step 2: Clear any lingering peft_config from the unwrapped model.
|
||||
# After model.unload(), the base model may still carry a peft_config
|
||||
# attribute. Removing it ensures PeftModel.from_pretrained() gets
|
||||
# a clean base model without "multiple adapters" warnings.
|
||||
if hasattr(model, 'peft_config'):
|
||||
del model.peft_config
|
||||
|
||||
logger.info("Model has been successfully reverted to a clean base state.")
|
||||
logger.info(f"Model '{base_model_name}' reverted to clean base state.")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -204,35 +201,29 @@ class InferenceBackend:
|
|||
import traceback
|
||||
logger.error(traceback.format_exc())
|
||||
return False
|
||||
pass
|
||||
|
||||
def activate_lora_adapter(self, base_model_name: str, lora_path: str) -> Tuple[bool, Optional[str]]:
|
||||
"""
|
||||
Activates a specific LoRA adapter on what is assumed to be a clean base model.
|
||||
Uses PeftModel.from_pretrained() which correctly wraps the base model.
|
||||
"""
|
||||
model = self.models[base_model_name].get("model")
|
||||
adapter_name_to_load = lora_path.split("/")[-1].replace(".", "_")
|
||||
|
||||
try:
|
||||
# At this point, the model should be clean thanks to revert_to_base_model.
|
||||
# We can now safely load and set the new adapter.
|
||||
|
||||
# Step 3: Load the new adapter.
|
||||
logger.info(f"Loading adapter '{adapter_name_to_load}' from '{lora_path}'")
|
||||
model.load_adapter(lora_path, adapter_name=adapter_name_to_load)
|
||||
|
||||
# Step 4: Set the new adapter as active.
|
||||
logger.info(f"Setting '{adapter_name_to_load}' as the active adapter.")
|
||||
model.set_adapter(adapter_name_to_load)
|
||||
# Use PeftModel.from_pretrained to wrap the clean base model with the adapter.
|
||||
# This is the correct approach after model.unload() + del peft_config.
|
||||
logger.info(f"Loading LoRA adapter '{adapter_name_to_load}' from '{lora_path}'...")
|
||||
model = PeftModel.from_pretrained(model, lora_path, adapter_name=adapter_name_to_load)
|
||||
self.models[base_model_name]["model"] = model
|
||||
logger.info(f"LoRA adapter '{adapter_name_to_load}' activated successfully.")
|
||||
|
||||
return True, adapter_name_to_load
|
||||
except Exception as e:
|
||||
# This will catch the "already exists" error if revert_to_base_model failed.
|
||||
logger.error(f"Failed to activate LoRA adapter '{adapter_name_to_load}': {e}")
|
||||
import traceback
|
||||
logger.error(traceback.format_exc())
|
||||
return False, None
|
||||
pass
|
||||
|
||||
def load_adapter(self, base_model_name: str, adapter_path: str, adapter_name: str = None) -> bool:
|
||||
"""
|
||||
|
|
@ -448,6 +439,74 @@ 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.
|
||||
|
||||
Uses revert_to_base_model() / activate_lora_adapter() which work correctly
|
||||
for models loaded by Unsloth as complete PeftModels (via model.unload() /
|
||||
model.load_adapter()), matching the proven pattern from the Gradio eval page.
|
||||
|
||||
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:
|
||||
# Revert to pure base model by unloading adapter weights
|
||||
logger.info(f"Compare mode: reverting '{base}' to base model for generation")
|
||||
self.revert_to_base_model(base)
|
||||
|
||||
elif use_adapter is True:
|
||||
# Activate the LoRA adapter from the original model path
|
||||
lora_path = model_info.get("model_path")
|
||||
if lora_path and model_info.get("is_lora"):
|
||||
logger.info(f"Compare mode: activating LoRA adapter from '{lora_path}' on '{base}'")
|
||||
self.activate_lora_adapter(base, lora_path)
|
||||
else:
|
||||
# Fallback for dynamically attached adapters
|
||||
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 adapter path/adapters on model")
|
||||
|
||||
elif isinstance(use_adapter, str):
|
||||
# Activate a specific adapter by path
|
||||
logger.info(f"Compare mode: activating specific adapter '{use_adapter}' on '{base}'")
|
||||
self.activate_lora_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 +518,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 +554,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 ────────────────────────────────────
|
||||
|
|
|
|||
|
|
@ -69,10 +69,10 @@ async def load_model(request: LoadRequest):
|
|||
backend = get_inference_backend()
|
||||
|
||||
# Create config using clean factory method
|
||||
# is_lora is auto-detected from adapter_config.json on disk/HF
|
||||
config = ModelConfig.from_identifier(
|
||||
model_id=request.model_path,
|
||||
hf_token=request.hf_token,
|
||||
is_lora=request.is_lora,
|
||||
)
|
||||
|
||||
if not config:
|
||||
|
|
@ -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