unsloth/studio/backend/core/inference/inference.py

1222 lines
49 KiB
Python

"""
Core inference backend - streamlined
"""
from unsloth import FastLanguageModel, FastVisionModel
from unsloth.chat_templates import get_chat_template
from transformers import TextStreamer
from peft import PeftModel, PeftModelForCausalLM
import sys
import torch
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
from utils.hardware import get_device, clear_gpu_cache, log_gpu_memory
from io import StringIO
import logging
logger = logging.getLogger(__name__)
class InferenceBackend:
"""Unified inference backend supporting text, vision, and LoRA models"""
def __init__(self):
self.models = {}
self.active_model_name = None
self.loading_models = set()
self.loaded_local_models = [] # [(display_name, path), ...]
self.default_models = [
"unsloth/Qwen3-4B-Instruct-2507",
"unsloth/Meta-Llama-3.1-8B-Instruct-bnb-4bit",
"unsloth/Mistral-Nemo-Instruct-2407-bnb-4bit",
"unsloth/Phi-3.5-mini-instruct",
"unsloth/Gemma-3-4B-it",
"unsloth/Qwen2-VL-2B-Instruct-bnb-4bit",
]
self.device = get_device().value
# Thread safety
import threading
self._generation_lock = threading.RLock()
self._model_state_lock = threading.Lock()
logger.info(f"InferenceBackend initialized on {self.device}")
def load_model(self,
config: ModelConfig,
max_seq_length: int = 2048,
dtype = None,
load_in_4bit: bool = True,
hf_token: Optional[str] = None) -> bool:
"""
Load any model: base, LoRA adapter, text, or vision.
"""
try:
model_name = config.identifier
# Check if already loaded
if model_name in self.models and self.models[model_name].get("model"):
logger.info(f"Model {model_name} already loaded")
self.active_model_name = model_name
return True
# Check if currently loading
if model_name in self.loading_models:
logger.info(f"Model {model_name} is already being loaded")
return False
self.loading_models.add(model_name)
self.models[model_name] = {
"is_vision": config.is_vision,
"is_lora": config.is_lora,
"model_path": config.path,
"base_model": config.base_model if config.is_lora else None,
"loaded_adapters": {},
"active_adapter": None,
}
model_type = "vision" if config.is_vision else "text"
adapter_info = " (LoRA adapter)" if self.models[model_name]["is_lora"] else ""
logger.info(f"Loading {model_type} model{adapter_info}: {model_name}")
log_gpu_memory(f"Before loading {model_name}")
# Load model - same approach for base models and LoRA adapters
if config.is_vision:
# Vision model (or vision LoRA adapter)
model, processor = FastVisionModel.from_pretrained(
model_name=config.path, # Can be base model OR LoRA adapter path
max_seq_length=max_seq_length,
dtype=dtype,
load_in_4bit=load_in_4bit,
token=hf_token if hf_token and hf_token.strip() else None,
)
# Apply inference optimization
FastVisionModel.for_inference(model)
self.models[model_name]["model"] = model
self.models[model_name]["tokenizer"] = processor
self.models[model_name]["processor"] = processor
else:
# Text model (or text LoRA adapter)
model, tokenizer = FastLanguageModel.from_pretrained(
model_name=config.path, # Can be base model OR LoRA adapter path
max_seq_length=max_seq_length,
dtype=dtype,
load_in_4bit=load_in_4bit,
token=hf_token if hf_token and hf_token.strip() else None,
)
# Apply inference optimization
FastLanguageModel.for_inference(model)
self.models[model_name]["model"] = model
self.models[model_name]["tokenizer"] = tokenizer
# Load chat template info
self._load_chat_template_info(model_name)
self.active_model_name = model_name
self.loading_models.discard(model_name)
logger.info(f"Successfully loaded model: {model_name}")
log_gpu_memory(f"After loading {model_name}")
return True
except Exception as e:
logger.error(f"Failed to load model: {e}")
error_msg = format_error_message(e, config.identifier)
# Cleanup on failure
if model_name in self.models:
del self.models[model_name]
self.loading_models.discard(model_name)
raise Exception(error_msg)
pass
# Add this new function
def unload_model(self, model_name: str) -> bool:
"""
Completely removes a model from the registry and clears GPU memory.
"""
if model_name in self.models:
try:
logger.info(f"Unloading model '{model_name}' from memory.")
# Delete the model entry from our registry
del self.models[model_name]
# Clear the active model if it was the one being unloaded
if self.active_model_name == model_name:
self.active_model_name = None
# Clear GPU memory cache
clear_gpu_cache()
logger.info(f"Model '{model_name}' successfully unloaded.")
return True
except Exception as e:
logger.error(f"Error while unloading model '{model_name}': {e}")
return False
else:
logger.warning(f"Attempted to unload model '{model_name}', but it was not found in the registry.")
return True
pass
def revert_to_base_model(self, base_model_name: str) -> bool:
"""
Reverts the model to its pristine base state by unloading AND
deleting all adapter configurations, as instructed.
"""
if base_model_name not in self.models:
return False
model = self.models[base_model_name].get("model")
print(f"[DEBUG] revert_to_base_model called. Model type BEFORE: {model.__class__.__name__}")
print(f"[DEBUG] Is PeftModel? {isinstance(model, (PeftModel, PeftModelForCausalLM))}")
print(f"[DEBUG] Has peft_config? {hasattr(model, 'peft_config')}, keys={list(getattr(model, 'peft_config', {}).keys())}")
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.
if isinstance(model, (PeftModel, PeftModelForCausalLM)):
print("[DEBUG] Model IS a PeftModel. Calling model.unload()...")
unwrapped_base_model = model.unload()
self.models[base_model_name]["model"] = unwrapped_base_model
model = unwrapped_base_model # Continue with the unwrapped model
print(f"[DEBUG] Model type AFTER unload: {model.__class__.__name__}")
else:
print(f"[DEBUG] Model is NOT a PeftModel, skipping unload.")
# 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:
print(f"[DEBUG] Lingering peft_config keys: {list(model.peft_config.keys())}")
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()):
if name == "default":
continue
logger.info(f"Deleting adapter config: '{name}'")
model.delete_adapter(name)
print(f"[DEBUG] Model type FINAL: {model.__class__.__name__}. Reverted to clean base state.")
return True
except Exception as e:
logger.error(f"Failed to revert model to base state: {e}")
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.
"""
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)
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:
"""
Load a LoRA adapter onto the base model if it's not already registered.
This method is idempotent.
"""
if base_model_name not in self.models:
logger.error(f"Base model {base_model_name} not loaded")
return False
model = self.models[base_model_name].get("model")
if model is None:
logger.error(f"Model object for {base_model_name} is None.")
return False
if adapter_name is None:
adapter_name = adapter_path.split("/")[-1].replace(".", "_")
# If we've loaded this adapter before, we don't need to do anything.
if adapter_name in self.models[base_model_name].get("loaded_adapters", {}):
logger.info(f"Adapter '{adapter_name}' is already registered. Skipping.")
return True
try:
logger.info(f"Loading new adapter '{adapter_name}' from '{adapter_path}' onto {base_model_name}")
# Unsloth modifies the model in-place and returns None. Do NOT re-assign.
model.load_adapter(adapter_path, adapter_name=adapter_name)
# Update our internal registry so we don't load it again.
self.models[base_model_name]["loaded_adapters"][adapter_name] = adapter_path
total_adapters = len(getattr(model, 'peft_config', {}))
logger.info(f"Adapter '{adapter_name}' loaded successfully. (Total adapters on model: {total_adapters})")
return True
except Exception as e:
logger.error(f"Failed to load adapter '{adapter_name}': {e}")
import traceback
logger.error(traceback.format_exc())
return False
pass
def enable_adapter(self, base_model_name: str, adapter_name: str) -> bool:
"""Enable specific adapter (for generation)"""
if base_model_name not in self.models:
return False
model = self.models[base_model_name]["model"]
try:
logger.info(f"Enabling adapter: {adapter_name}")
model.set_adapter(adapter_name)
self.models[base_model_name]["active_adapter"] = adapter_name
return True
except Exception as e:
logger.error(f"Failed to enable adapter: {e}")
return False
def disable_adapters(self, base_model_name: str) -> bool:
"""Disable all adapters (back to pure base model)"""
if base_model_name not in self.models:
return False
model = self.models[base_model_name]["model"]
try:
logger.info(f"Disabling all adapters on {base_model_name}")
model.disable_adapters()
self.models[base_model_name]["active_adapter"] = None
return True
except Exception as e:
logger.error(f"Failed to disable adapters: {e}")
return False
# In backend/inference.py
def load_for_eval(self, lora_path: str, max_seq_length: int = 2048,
dtype = None, load_in_4bit: bool = True,
hf_token: Optional[str] = None) -> Tuple[bool, Optional[str], Optional[str]]:
"""
Prepare for eval: ensure base model and the specified adapter are loaded.
"""
try:
from utils.models import ModelConfig
lora_config = ModelConfig.from_lora_path(lora_path, hf_token)
if not lora_config:
return False, None, None
base_model_name = lora_config.base_model
# 1. Load the base model if it's not already in memory (this logic is correct)
if base_model_name not in self.models or not self.models[base_model_name].get("model"):
logger.info(f"Base model '{base_model_name}' not loaded, loading now.")
base_config = ModelConfig.from_ui_selection(base_model_name, None, is_lora=False)
if not self.load_model(base_config, max_seq_length, dtype, load_in_4bit, hf_token):
return False, None, None
else:
logger.info(f"Base model '{base_model_name}' is already in memory.")
self.active_model_name = base_model_name
# 2. Delegate to our now-idempotent load_adapter function.
# It will handle all cases: first adapter, or subsequent adapters.
adapter_name = lora_path.split("/")[-1].replace(".", "_")
adapter_success = self.load_adapter(
base_model_name=base_model_name,
adapter_path=lora_path,
adapter_name=adapter_name
)
if not adapter_success:
return False, base_model_name, None
return True, base_model_name, adapter_name
except Exception as e:
logger.error(f"Error during load_for_eval: {e}")
import traceback
logger.error(traceback.format_exc())
return False, None, None
pass
def load_for_eval(self, lora_path: str, max_seq_length: int = 2048,
dtype = None, load_in_4bit: bool = True,
hf_token: Optional[str] = None) -> Tuple[bool, Optional[str], Optional[str]]:
"""
Final Corrected Version:
Ensures the base model and the specified adapter are loaded.
This function is idempotent and handles all states correctly.
"""
try:
from utils.models import ModelConfig
lora_config = ModelConfig.from_lora_path(lora_path, hf_token)
if not lora_config:
return False, None, None
base_model_name = lora_config.base_model
# 1. Load the base model if it's not already in memory
if base_model_name not in self.models or not self.models[base_model_name].get("model"):
logger.info(f"Base model '{base_model_name}' not loaded, loading now.")
base_config = ModelConfig.from_ui_selection(base_model_name, None, is_lora=False)
if not self.load_model(base_config, max_seq_length, dtype, load_in_4bit, hf_token):
return False, None, None
self.active_model_name = base_model_name
# 2. Determine the required adapter name from the user's selection
adapter_name = lora_path.split("/")[-1].replace(".", "_")
# 3. Call our robust load_adapter function to ensure this specific adapter is loaded.
# It will only load from disk if the model doesn't already have it.
adapter_success = self.load_adapter(
base_model_name=base_model_name,
adapter_path=lora_path,
adapter_name=adapter_name
)
if not adapter_success:
return False, base_model_name, None
# 4. Return the correct, verified adapter name for the UI logic to use.
return True, base_model_name, adapter_name
except Exception as e:
logger.error(f"Error during load_for_eval: {e}")
import traceback
logger.error(traceback.format_exc())
return False, None, None
pass
def load_adapter(self, base_model_name: str, adapter_path: str, adapter_name: str) -> bool:
"""
Loads an adapter onto the model ONLY if it's not already attached.
"""
model = self.models[base_model_name].get("model")
# Check if this adapter name is already part of the model's config. This is the most reliable check.
if hasattr(model, "peft_config") and adapter_name in model.peft_config:
logger.info(f"Adapter '{adapter_name}' is already attached to the model. Skipping load.")
return True
try:
logger.info(f"Loading new adapter '{adapter_name}' from '{adapter_path}' onto {base_model_name}")
model.load_adapter(adapter_path, adapter_name=adapter_name)
# Update our internal registry ONLY after a successful load.
if "loaded_adapters" not in self.models[base_model_name]:
self.models[base_model_name]["loaded_adapters"] = {}
self.models[base_model_name]["loaded_adapters"][adapter_name] = adapter_path
total_adapters = len(getattr(model, 'peft_config', {}))
logger.info(f"Adapter '{adapter_name}' loaded successfully. (Total unique adapters on model: {total_adapters})")
return True
except Exception as e:
logger.error(f"Failed to load adapter '{adapter_name}': {e}")
return False
pass
def set_active_adapter(self, base_model_name: str, adapter_name: str) -> bool:
"""
Sets the active adapter for generation. This replaces the flawed 'enable_adapter'.
"""
model = self.models[base_model_name].get("model")
try:
logger.info(f"Setting active adapter to: '{adapter_name}'")
model.set_adapter(adapter_name)
self.models[base_model_name]["active_adapter"] = adapter_name
return True
except Exception as e:
# This will catch the "adapter not found" error if something goes wrong.
logger.error(f"Failed to set active adapter to '{adapter_name}': {e}")
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,
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]:
"""
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,
)
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"
return
model_info = self.models[self.active_model_name]
is_vision = model_info.get("is_vision", False)
tokenizer = model_info.get("tokenizer") or model_info.get("processor")
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
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}")
# 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
)
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,
repetition_penalty) -> Generator[str, None, None]:
"""Handle vision model generation with true token-by-token streaming."""
model_info = self.models[self.active_model_name]
model = model_info["model"]
processor = model_info["processor"]
# Extract user message
user_message = ""
if messages and messages[-1]["role"] == "user":
import re
user_message = messages[-1]["content"]
user_message = re.sub(r'<img[^>]*>', '', user_message).strip()
if not user_message:
user_message = "Describe this image." if image else "Hello"
# Prepare vision messages
if image:
vision_messages = [
{
"role": "user",
"content": [
{"type": "image"},
{"type": "text", "text": user_message}
],
}
]
input_text = processor.apply_chat_template(vision_messages, add_generation_prompt=True)
inputs = processor(
image,
input_text,
add_special_tokens=False,
return_tensors="pt",
).to(self.device)
else:
# Text-only for vision model
formatted_prompt = self.format_chat_prompt(messages, system_prompt)
inputs = processor.tokenizer(formatted_prompt, return_tensors="pt").to(self.device)
# Stream with TextIteratorStreamer + background thread
try:
from transformers import TextIteratorStreamer
import threading
streamer = TextIteratorStreamer(
processor.tokenizer, skip_prompt=True, skip_special_tokens=True
)
generation_kwargs = dict(
**inputs,
streamer=streamer,
max_new_tokens=max_new_tokens,
use_cache=True,
temperature=temperature,
top_p=top_p,
top_k=top_k,
)
def generate_fn():
try:
model.generate(**generation_kwargs)
except Exception as e:
logger.error(f"Vision generation error in thread: {e}")
thread = threading.Thread(target=generate_fn)
thread.start()
output = ""
for new_token in streamer:
if new_token:
output += new_token
cleaned = self._clean_generated_text(output)
yield cleaned
thread.join()
except Exception as e:
logger.error(f"Vision generation error: {e}")
yield f"Error: {str(e)}"
pass
def generate_stream(self,
prompt: str,
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]:
"""Generate streaming text response (text models only)."""
if not self.active_model_name:
yield "Error: No active model"
return
model_info = self.models[self.active_model_name]
model = model_info["model"]
tokenizer = model_info["tokenizer"]
try:
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
from transformers import TextIteratorStreamer
import threading
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
generation_kwargs = dict(
**inputs,
streamer=streamer,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
repetition_penalty=repetition_penalty,
do_sample=True,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.eos_token_id if tokenizer.pad_token_id is None else tokenizer.pad_token_id,
)
def generate_fn():
try:
model.generate(**generation_kwargs)
except Exception as e:
logger.error(f"Generation error: {e}")
thread = threading.Thread(target=generate_fn)
thread.start()
output = ""
for new_token in streamer:
if new_token:
output += new_token
cleaned = self._clean_generated_text(output)
yield cleaned
thread.join()
except Exception as e:
logger.error(f"Error during generation: {e}")
yield f"Error: {str(e)}"
# ... other helper methods (format_chat_prompt, _clean_generated_text, etc.)
pass
def format_chat_prompt(self, messages: list, system_prompt: str = None) -> str:
if not self.active_model_name or self.active_model_name not in self.models:
logger.error("No active model available")
return ""
if self.models[self.active_model_name].get("tokenizer") is None:
logger.error("Tokenizer not loaded for active model")
return ""
chat_template_info = self.models[self.active_model_name].get("chat_template_info", {})
tokenizer = self.models[self.active_model_name]["tokenizer"]
chat_messages = []
if system_prompt:
chat_messages.append({"role": "system", "content": system_prompt})
last_role = "system" if system_prompt else None
for msg in messages:
role = msg.get("role", "")
content = msg.get("content", "")
if role in ["system", "user", "assistant"] and content.strip():
if role == last_role:
logger.debug(f"Skipping consecutive {role} message to maintain alternation")
continue
if role == "user":
import re
clean_content = re.sub(r'<[^>]+>', '', content).strip()
if clean_content:
chat_messages.append({"role": role, "content": clean_content})
last_role = role
elif role == "assistant" and content.strip():
chat_messages.append({"role": role, "content": content})
last_role = role
elif role == "system":
continue
if chat_messages and chat_messages[-1]["role"] == "assistant":
logger.debug("Removing final assistant message to ensure proper alternation")
chat_messages.pop()
logger.info(f"Sending {len(chat_messages)} messages to tokenizer:")
for i, msg in enumerate(chat_messages):
logger.info(f" {i}: {msg['role']} - {msg['content'][:50]}...")
try:
formatted_prompt = tokenizer.apply_chat_template(
chat_messages,
tokenize=False,
add_generation_prompt=True
)
logger.info(f"Successfully applied tokenizer's native chat template")
return formatted_prompt
except Exception as e:
error_msg = str(e).lower()
if "chat_template is not set" in error_msg or "no template argument" in error_msg:
logger.info(f"Base model detected - no built-in chat template available, using fallback formatting")
else:
logger.warning(f"Failed to apply tokenizer chat template: {e}")
logger.debug(f"""Failed with messages: {[f"{m['role']}: {m['content'][:30]}..." for m in chat_messages]}""")
if chat_template_info.get("has_template", False):
logger.info("Falling back to manual template formatting based on detected patterns")
template_type = chat_template_info.get("format_type", "generic")
manual_prompt = self._format_chat_manual(chat_messages, template_type, chat_template_info.get("special_tokens", {}))
logger.info(f"Manual template result: {manual_prompt[:200]}...")
return manual_prompt
else:
logger.info("Using generic chat formatting for base model")
return self._format_generic_template(chat_messages, {})
def _format_chat_manual(self, messages: list, template_type: str, special_tokens: dict) -> str:
"""
Manual chat formatting fallback for when tokenizer template fails
Args:
messages: List of message dictionaries
template_type: Detected template type
special_tokens: Dictionary of special tokens
Returns:
str: Manually formatted prompt
"""
if template_type == "llama3":
return self._format_llama3_template(messages, special_tokens)
elif template_type == "mistral":
return self._format_mistral_template(messages, special_tokens)
elif template_type == "chatml":
return self._format_chatml_template(messages, special_tokens)
elif template_type == "alpaca":
return self._format_alpaca_template(messages, special_tokens)
else:
return self._format_generic_template(messages, special_tokens)
def _format_llama3_template(self, messages: list, special_tokens: dict) -> str:
"""Format messages using Llama 3 template"""
bos_token = special_tokens.get("bos_token", "<|begin_of_text|>")
formatted = bos_token
for msg in messages:
role = msg["role"]
content = msg["content"]
formatted += f"<|start_header_id|>{role}<|end_header_id|>\n\n{content}<|eot_id|>"
formatted += "<|start_header_id|>assistant<|end_header_id|>\n\n"
return formatted
def _format_mistral_template(self, messages: list, special_tokens: dict) -> str:
"""Format messages using Mistral template"""
bos_token = special_tokens.get("bos_token", "<s>")
formatted = bos_token
system_msg = None
conversation = []
for msg in messages:
if msg["role"] == "system":
system_msg = msg["content"]
else:
conversation.append(msg)
i = 0
while i < len(conversation):
if conversation[i]["role"] == "user":
user_content = conversation[i]["content"]
if system_msg and i == 0:
user_content = f"{system_msg}\n\n{user_content}"
formatted += f"[INST] {user_content} [/INST]"
if i + 1 < len(conversation) and conversation[i + 1]["role"] == "assistant":
formatted += f" {conversation[i + 1]['content']}</s>"
i += 2
else:
formatted += " "
break
else:
i += 1
return formatted
def _format_chatml_template(self, messages: list, special_tokens: dict) -> str:
"""Format messages using ChatML template"""
formatted = ""
for msg in messages:
role = msg["role"]
content = msg["content"]
formatted += f"<|im_start|>{role}\n{content}<|im_end|>\n"
formatted += "<|im_start|>assistant\n"
return formatted
def _format_alpaca_template(self, messages: list, special_tokens: dict) -> str:
"""Format messages using Alpaca template"""
formatted = ""
system_msg = None
for msg in messages:
if msg["role"] == "system":
system_msg = msg["content"]
elif msg["role"] == "user":
if system_msg:
formatted += f"### Instruction:\n{system_msg}\n\n### Input:\n{msg['content']}\n\n### Response:\n"
system_msg = None
else:
formatted += f"### Human:\n{msg['content']}\n\n### Assistant:\n"
elif msg["role"] == "assistant":
formatted += f"{msg['content']}\n\n"
return formatted
def _format_generic_template(self, messages: list, special_tokens: dict) -> str:
"""Generic fallback formatting"""
formatted = ""
for msg in messages:
role = msg["role"].title()
content = msg["content"]
formatted += f"{role}: {content}\n"
formatted += "Assistant: "
return formatted
def check_vision_model_compatibility(self) -> bool:
"""
Check if current model supports vision.
Returns:
bool: True if current model supports vision, False otherwise
"""
current_model = self.get_current_model()
if current_model and current_model in self.models:
return self.models[current_model].get("is_vision", False)
return False
def _reset_model_generation_state(self, model_name: str):
"""Reset generation state for a specific model to prevent contamination."""
if model_name not in self.models:
return
model = self.models[model_name].get("model")
if not model:
return
try:
# This is a common pattern for Unsloth/Hugging Face models
if hasattr(model, 'past_key_values'):
model.past_key_values = None
if hasattr(model, 'generation_config'):
if hasattr(model.generation_config, 'past_key_values'):
model.generation_config.past_key_values = None
logger.debug(f"Reset generation state for model: {model_name}")
except Exception as e:
logger.warning(f"Could not fully reset model state for {model_name}: {e}")
pass
def reset_generation_state(self):
"""Reset any cached generation state to prevent hanging after errors"""
try:
# Clear cached states for ALL loaded models
for model_name in self.models.keys():
self._reset_model_generation_state(model_name)
clear_gpu_cache()
logger.debug("Cleared GPU cache")
import gc
gc.collect()
logger.info("Performed comprehensive generation state reset")
except Exception as e:
logger.warning(f"Could not fully reset generation state: {e}")
def resize_image(self, img, max_size: int = 800):
"""Resize image while maintaining aspect ratio if either dimension exceeds max_size"""
if img is None:
return None
if img.size[0] > max_size or img.size[1] > max_size:
from PIL import Image
ratio = min(max_size/img.size[0], max_size/img.size[1])
new_size = (int(img.size[0]*ratio), int(img.size[1]*ratio))
return img.resize(new_size, Image.Resampling.LANCZOS)
return img
def _clean_generated_text(self, text: str) -> str:
import re
text = re.sub(r'<\|start_header_id\|>.*?<\|end_header_id\|>', '', text)
text = re.sub(r'<\|eot_id\|>', '', text)
text = re.sub(r'<\|begin_of_text\|>', '', text)
text = re.sub(r'\[INST\].*?\[/INST\]', '', text)
text = re.sub(r'<s>|</s>', '', text)
# Clean ChatML tokens (used by Qwen2-VL and similar models)
text = re.sub(r'<\|im_start\|>.*?<\|im_end\|>', '', text)
text = re.sub(r'<\|im_end\|>', '', text)
text = re.sub(r'<\|im_start\|>', '', text)
text = re.sub(r'^\s*(assistant|user|system):\s*', '', text, flags=re.IGNORECASE)
text = text.strip()
return text
def _load_chat_template_info(self, model_name: str):
if model_name not in self.models or not self.models[model_name].get("tokenizer"):
return
tokenizer = self.models[model_name]["tokenizer"]
chat_template_info = {
"has_template": False,
"template": None,
"format_type": "generic",
"special_tokens": {},
"template_name": None,
}
try:
from utils.datasets import MODEL_TO_TEMPLATE_MAPPER
#Try exact match first
model_name_lower = model_name.lower()
if model_name_lower in MODEL_TO_TEMPLATE_MAPPER:
chat_template_info["template_name"] = MODEL_TO_TEMPLATE_MAPPER[model_name_lower]
logger.info(f"Detected template '{chat_template_info['template_name']}' for {model_name} from mapper")
else:
# Try partial match (for variants like model_name-bnb-4bit)
for key in MODEL_TO_TEMPLATE_MAPPER:
if key in model_name_lower or model_name_lower in key:
chat_template_info["template_name"] = MODEL_TO_TEMPLATE_MAPPER[key]
logger.info(f"Detected template '{chat_template_info['template_name']}' for {model_name} (partial match)")
break
except Exception as e:
logger.warning(f"Could not detect template from mapper for {model_name}: {e}")
try:
if hasattr(tokenizer, 'chat_template') and tokenizer.chat_template:
chat_template_info["has_template"] = True
chat_template_info["template"] = tokenizer.chat_template
template_str = tokenizer.chat_template.lower()
if "start_header_id" in template_str and "end_header_id" in template_str:
chat_template_info["format_type"] = "llama3"
elif "[inst]" in template_str and "[/inst]" in template_str:
chat_template_info["format_type"] = "mistral"
elif "<|im_start|>" in template_str and "<|im_end|>" in template_str:
chat_template_info["format_type"] = "chatml"
elif "### instruction:" in template_str or "### human:" in template_str:
chat_template_info["format_type"] = "alpaca"
else:
chat_template_info["format_type"] = "custom"
logger.info(f"Loaded chat template for {model_name} (detected as {chat_template_info['format_type']} format)")
logger.debug(f"Template preview: {tokenizer.chat_template[:200]}...")
special_tokens = {}
if hasattr(tokenizer, 'bos_token') and tokenizer.bos_token:
special_tokens["bos_token"] = tokenizer.bos_token
if hasattr(tokenizer, 'eos_token') and tokenizer.eos_token:
special_tokens["eos_token"] = tokenizer.eos_token
if hasattr(tokenizer, 'pad_token') and tokenizer.pad_token:
special_tokens["pad_token"] = tokenizer.pad_token
chat_template_info["special_tokens"] = special_tokens
else:
logger.info(f"No chat template found for {model_name}, will use generic formatting")
except Exception as e:
logger.error(f"Error loading chat template info for {model_name}: {e}")
self.models[model_name]["chat_template_info"] = chat_template_info
if chat_template_info["has_template"]:
logger.info(f"Chat template loaded for {model_name}: {chat_template_info['format_type']} format")
else:
logger.info(f"No built-in chat template for {model_name}, will use generic formatting")
def get_current_model(self) -> Optional[str]:
"""Get currently active model name"""
return self.active_model_name
def is_model_loading(self) -> bool:
"""Check if any model is currently loading"""
return len(self.loading_models) > 0
def get_loading_model(self) -> Optional[str]:
"""Get name of currently loading model"""
return next(iter(self.loading_models)) if self.loading_models else None
def load_model_simple(self,
model_path: str,
hf_token: Optional[str] = None,
max_seq_length: int = 2048,
load_in_4bit: bool = True) -> bool:
"""
Simple model loading wrapper for chat interface.
Accepts model path as string and handles ModelConfig creation internally.
Args:
model_path: Model name or path (e.g., "unsloth/llama-3-8b")
hf_token: HuggingFace token for gated models
max_seq_length: Maximum sequence length
load_in_4bit: Whether to use 4-bit quantization
Returns:
bool: True if successful, False otherwise
"""
try:
# Create config from string path
config = ModelConfig.from_ui_selection(
model_path,
lora_path=None, # No LoRA for chat
is_lora=False
)
# Call existing load_model with config
return self.load_model(
config=config,
max_seq_length=max_seq_length,
dtype=None, # Auto-detect
load_in_4bit=load_in_4bit,
hf_token=hf_token
)
except Exception as e:
logger.error(f"Error in load_model_simple: {e}")
return False
def load_model_simple(self,
model_path: str,
hf_token: Optional[str] = None,
max_seq_length: int = 2048,
load_in_4bit: bool = True) -> bool:
"""
Simple model loading wrapper for chat interface.
Accepts model path as string and handles ModelConfig creation internally.
Args:
model_path: Model name or path (e.g., "unsloth/llama-3-8b")
hf_token: HuggingFace token for gated models
max_seq_length: Maximum sequence length
load_in_4bit: Whether to use 4-bit quantization
Returns:
bool: True if successful, False otherwise
"""
try:
from backend.model_config import ModelConfig
logger.info(f"load_model_simple called with: {model_path}")
# Create config from string path
config = ModelConfig.from_ui_selection(
model_path,
lora_path=None, # No LoRA for chat
is_lora=False
)
logger.info(f"Created ModelConfig with identifier: {config.identifier}")
# Call existing load_model with config
return self.load_model(
config=config,
max_seq_length=max_seq_length,
dtype=None, # Auto-detect
load_in_4bit=load_in_4bit,
hf_token=hf_token
)
except Exception as e:
logger.error(f"Error in load_model_simple: {e}")
import traceback
traceback.print_exc()
return False
pass
# Global inference backend instance
inference_backend = InferenceBackend()
def get_inference_backend() -> InferenceBackend:
return inference_backend