""" 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 json import sys import torch from pathlib import Path 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 — _generation_lock serializes model.generate() calls. # Must be a regular Lock (NOT RLock) because in async FastAPI, multiple # requests share the same event-loop thread, so RLock reentrancy lets # concurrent compare-mode requests race on the GPU. The lock is # acquired by the *background generation thread*, not the event-loop. import threading self._generation_lock = threading.Lock() self._model_state_lock = threading.Lock() logger.info(f"InferenceBackend initialized on {self.device}") @staticmethod def _normalize_top_k(top_k: int) -> int: # API supports -1 as "disable top-k"; transformers expects 0 to disable. return 0 if top_k < 0 else top_k 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) # FastVisionModel may return a raw tokenizer (e.g. GemmaTokenizerFast) # instead of a proper Processor for some models (e.g. Gemma-3). # In that case, load the real processor from the base model. from transformers import ProcessorMixin if not (isinstance(processor, ProcessorMixin) or hasattr(processor, "image_processor")): # For LoRA adapters, use the base model. For local merged exports, # read export_metadata.json to find the original base model. processor_source = config.base_model if config.is_lora else config.identifier if not config.is_lora and config.is_local: _meta_path = Path(config.path) / "export_metadata.json" try: if _meta_path.exists(): _meta = json.loads(_meta_path.read_text()) if _meta.get("base_model"): processor_source = _meta["base_model"] except Exception: pass logger.warning( f"FastVisionModel returned {type(processor).__name__} (no image_processor) " f"for '{model_name}' — loading proper processor from '{processor_source}'" ) from transformers import AutoProcessor processor = AutoProcessor.from_pretrained( processor_source, token=hf_token if hf_token and hf_token.strip() else None, ) logger.info(f"Loaded {type(processor).__name__} from {processor_source}") 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() # Remove stale compiled cache so the next model gets a fresh one from utils.cache_cleanup import clear_unsloth_compiled_cache clear_unsloth_compiled_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") try: # Step 1: Unload the adapter weights if model is a PeftModel. if isinstance(model, (PeftModel, PeftModelForCausalLM)): 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 # 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(f"Model '{base_model_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 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: # 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: logger.error(f"Failed to activate LoRA adapter '{adapter_name_to_load}': {e}") import traceback logger.error(traceback.format_exc()) return False, None 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 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 PEFT's disable_adapter_layers() / enable_adapter_layers() which toggle a boolean flag on each LoRA layer. Unsloth's fast_linear_forward checks this flag (proj.disable_adapters) and skips LoRA computation when True. This is non-destructive — no model unloading/reloading needed. 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] model = model_info.get("model") if model is None: return if use_adapter is False: # Disable LoRA layers → base model output if isinstance(model, (PeftModel, PeftModelForCausalLM)): logger.info(f"Compare mode: disabling adapters on '{base}' for base model generation") model.base_model.disable_adapter_layers() else: logger.info(f"Compare mode: model '{base}' is not a PeftModel, already base") elif use_adapter is True: # Re-enable LoRA layers → adapter output if isinstance(model, (PeftModel, PeftModelForCausalLM)): logger.info(f"Compare mode: enabling adapters on '{base}' for LoRA generation") model.base_model.enable_adapter_layers() else: logger.warning("use_adapter=true but model is not a PeftModel") elif isinstance(use_adapter, str): # Enable adapters and set the specific one active if isinstance(model, (PeftModel, PeftModelForCausalLM)): logger.info(f"Compare mode: enabling adapter '{use_adapter}' on '{base}'") model.base_model.enable_adapter_layers() self.set_active_adapter(base, use_adapter) else: logger.warning(f"use_adapter='{use_adapter}' but model is not a PeftModel") def generate_with_adapter_control( self, use_adapter: Optional[Union[bool, str]] = None, cancel_event=None, **gen_kwargs, ) -> Generator[str, None, None]: """ Thread-safe generation with optional adapter toggling. The adapter toggle + model.generate() are serialized by _generation_lock inside the background generation thread — NOT in the event-loop thread. This prevents the RLock-reentrant race that occurs when two async SSE handlers share the same event-loop thread. Args: use_adapter: Adapter control (None/False/True/str). See _apply_adapter_state. **gen_kwargs: Forwarded to generate_chat_response. """ yield from self._generate_chat_response_inner( cancel_event=cancel_event, _adapter_state=use_adapter, **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, min_p: float = 0.0, max_new_tokens: int = 256, repetition_penalty: float = 1.1, cancel_event=None) -> Generator[str, None, None]: """ Generate response for text or vision models. The generation lock is acquired by the background generation thread. """ 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, min_p=min_p, max_new_tokens=max_new_tokens, repetition_penalty=repetition_penalty, cancel_event=cancel_event, ) 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, min_p: float = 0.0, max_new_tokens: int = 256, repetition_penalty: float = 1.1, cancel_event=None, _adapter_state=None) -> Generator[str, None, None]: """ Inner generation logic. Called by both generate_chat_response and generate_with_adapter_control. _adapter_state is passed to generate_stream/vision so the background thread can toggle adapters under the generation lock. """ 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") # Unwrap processor → raw tokenizer for VLMs on the text path tokenizer = getattr(tokenizer, "tokenizer", tokenizer) top_k = self._normalize_top_k(top_k) if is_vision and image: # Vision model generation (only when an image is actually provided) # Check that the stored processor can actually handle images. # FastVisionModel may return a raw tokenizer (e.g. GemmaTokenizerFast) # instead of a proper ProcessorMixin for some models (e.g. Gemma-3). from transformers import ProcessorMixin processor = model_info.get("processor") has_image_processing = ( processor is not None and (isinstance(processor, ProcessorMixin) or hasattr(processor, "image_processor")) ) if has_image_processing: yield from self._generate_vision_response( messages, system_prompt, image, temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty, cancel_event=cancel_event, ) return else: logger.warning( f"Model '{self.active_model_name}' is marked as vision but its processor " f"({type(processor).__name__}) has no image_processor — " f"falling back to text-only generation (image will be ignored)." ) # Text path: 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, chat_template=template_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, min_p, max_new_tokens, repetition_penalty, cancel_event=cancel_event, _adapter_state=_adapter_state, ) def _generate_vision_response(self, messages, system_prompt, image, temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty, cancel_event=None) -> 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"] # FastVisionModel may return a raw tokenizer (e.g. GemmaTokenizerFast) # instead of a Processor for some models. Safe unwrap for tokenize-only ops. raw_tokenizer = getattr(processor, "tokenizer", 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']*>', '', 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, tokenize=False) 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 = raw_tokenizer(formatted_prompt, return_tensors="pt").to(self.device) # Stream with TextIteratorStreamer + background thread try: from transformers import TextIteratorStreamer import threading streamer = TextIteratorStreamer( raw_tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=0.2, ) generation_kwargs = dict( **inputs, streamer=streamer, max_new_tokens=max_new_tokens, use_cache=True, do_sample=temperature > 0, temperature=temperature, top_p=top_p, top_k=top_k, min_p=min_p, ) err: dict[str, str] = {} def generate_fn(): with self._generation_lock: try: model.generate(**generation_kwargs) except Exception as e: err["msg"] = str(e) logger.error(f"Vision generation error in thread: {e}") finally: try: streamer.end() except Exception: pass thread = threading.Thread(target=generate_fn) thread.start() output = "" from queue import Empty try: while True: if cancel_event is not None and cancel_event.is_set(): break try: new_token = next(streamer) except StopIteration: break except Empty: if not thread.is_alive(): break continue if new_token: output += new_token cleaned = self._clean_generated_text(output) yield cleaned finally: if cancel_event is not None: cancel_event.set() thread.join(timeout=10) if thread.is_alive(): logger.warning("Vision generation thread did not exit after cancel/join timeout") if err.get("msg"): yield f"Error: {err['msg']}" 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, min_p: float = 0.0, max_new_tokens: int = 256, repetition_penalty: float = 1.1, cancel_event=None, _adapter_state=None) -> Generator[str, None, None]: """Generate streaming text response (text models only). _adapter_state: if not None, the background thread toggles adapters before model.generate(), all under _generation_lock. """ if not self.active_model_name: yield "Error: No active model" return model_info = self.models[self.active_model_name] model = model_info["model"] # For VLMs the stored "tokenizer" is actually the processor. # Unwrap to get the real tokenizer so TextIteratorStreamer's # skip_prompt / skip_special_tokens work correctly. tokenizer = model_info["tokenizer"] tokenizer = getattr(tokenizer, "tokenizer", 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, timeout=0.2, ) generation_kwargs = dict( **inputs, streamer=streamer, max_new_tokens=max_new_tokens, temperature=temperature, top_p=top_p, top_k=top_k, min_p=min_p, repetition_penalty=repetition_penalty, do_sample=temperature > 0, 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, ) if cancel_event is not None: from transformers.generation.stopping_criteria import ( StoppingCriteria, StoppingCriteriaList, ) class _CancelCriteria(StoppingCriteria): def __init__(self, ev): self.ev = ev def __call__(self, input_ids, scores, **kwargs): return self.ev.is_set() generation_kwargs["stopping_criteria"] = StoppingCriteriaList( [_CancelCriteria(cancel_event)] ) def generate_fn(): with self._generation_lock: try: if _adapter_state is not None: self._apply_adapter_state(_adapter_state) model.generate(**generation_kwargs) except Exception as e: err["msg"] = str(e) logger.error(f"Generation error: {e}") finally: try: streamer.end() except Exception: pass err: dict[str, str] = {} thread = threading.Thread(target=generate_fn) thread.start() output = "" from queue import Empty try: while True: if cancel_event is not None and cancel_event.is_set(): break try: new_token = next(streamer) except StopIteration: break except Empty: if not thread.is_alive(): break continue if new_token: output += new_token cleaned = self._clean_generated_text(output) yield cleaned finally: if cancel_event is not None: cancel_event.set() thread.join(timeout=10) if thread.is_alive(): logger.warning("Generation thread did not exit after cancel/join timeout") if err.get("msg"): yield f"Error: {err['msg']}" 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"] tokenizer = getattr(tokenizer, "tokenizer", 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", "") 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']}" 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: """Strip leaked special tokens using the tokenizer's own token list.""" tokenizer = self.models.get(self.active_model_name, {}).get("tokenizer") if tokenizer: for token in getattr(tokenizer, "all_special_tokens", []): if token in text: text = text.replace(token, "") return text.strip() 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: 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