diff --git a/studio/backend/backend/__init__.py b/studio/backend/backend/__init__.py deleted file mode 100644 index 33a28c425c..0000000000 --- a/studio/backend/backend/__init__.py +++ /dev/null @@ -1,44 +0,0 @@ -""" -Unified backend module for Unsloth -""" - -# Inference -from .inference import InferenceBackend - -# Training -from .trainer import UnslothTrainer, get_trainer -from .training import TrainingBackend, get_training_backend, create_training_handlers - -# Configuration -from .model_config import is_vision_model, ModelConfig, scan_trained_loras -# Utilities -from .path_utils import normalize_path, is_local_path, is_model_cached -from utils.utils import without_hf_auth, format_error_message, get_gpu_memory_info, search_hf_models -from utils.datasets.dataset_utils import format_and_template_dataset - -__all__ = [ - # Inference - 'InferenceBackend', - - # Training - 'UnslothTrainer', - 'get_trainer', - 'get_training_backend', - 'TrainingBackend', - "create_training_handlers", - - # Config - 'ModelConfig', - 'is_vision_model', - 'scan_trained_loras', - - # Utils - 'search_hf_models', - 'format_and_template_dataset', - 'normalize_path', - 'is_local_path', - 'is_model_cached', - 'without_hf_auth', - 'format_error_message', - 'get_gpu_memory_info', -] diff --git a/studio/backend/backend/export.py b/studio/backend/backend/export.py deleted file mode 100644 index b1557624ab..0000000000 --- a/studio/backend/backend/export.py +++ /dev/null @@ -1,506 +0,0 @@ -# backend/export.py -""" -Export backend - handles model exporting in various formats -""" -import logging -import os -from pathlib import Path -from typing import Optional, Tuple, List -from peft import PeftModel, PeftModelForCausalLM -from unsloth import FastLanguageModel, FastVisionModel -from huggingface_hub import HfApi, ModelCard -from transformers.modeling_utils import PushToHubMixin -import torch - -from .model_config import is_vision_model, get_base_model_from_lora -from .inference import get_inference_backend - -logger = logging.getLogger(__name__) - -# Model card template -MODEL_CARD = \ -"""--- -base_model: {base_model} -tags: -- text-generation-inference -- transformers -- unsloth -- {model_type} -- {extra} -license: apache-2.0 -language: -- en ---- - -# Uploaded finetuned {method} model - -- **Developed by:** {username} -- **License:** apache-2.0 -- **Finetuned from model :** {base_model} - -This {model_type} model was trained 2x faster with [Unsloth](https://github.com/unslothai/unsloth) and Huggingface's TRL library. - -[](https://github.com/unslothai/unsloth) -""" - -class ExportBackend: - """Handles model export operations""" - - def __init__(self): - self.inference_backend = get_inference_backend() - self.current_checkpoint = None - self.current_model = None - self.current_tokenizer = None - self.is_vision = False - self.is_peft = False - - def cleanup_memory(self): - """Offload and delete all models from memory""" - try: - logger.info("Starting memory cleanup...") - - # Unload all models from inference backend - model_names = list(self.inference_backend.models.keys()) - for model_name in model_names: - self.inference_backend.unload_model(model_name) - - # Clear current export state - self.current_model = None - self.current_tokenizer = None - self.current_checkpoint = None - - # Force garbage collection - import gc - gc.collect() - - # Clear CUDA cache - if torch.cuda.is_available(): - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - logger.info("Memory cleanup completed successfully") - return True - - except Exception as e: - logger.error(f"Error during memory cleanup: {e}") - return False - - def scan_checkpoints(self, outputs_dir: str = "./outputs") -> List[Tuple[str, str]]: - """ - Scan outputs folder for model checkpoints. - - Returns: - List of tuples: [(display_name, checkpoint_path), ...] - """ - checkpoints = [] - outputs_path = Path(outputs_dir) - - if not outputs_path.exists(): - logger.warning(f"Outputs directory not found: {outputs_dir}") - return checkpoints - - try: - for item in outputs_path.iterdir(): - if item.is_dir(): - # Check if this directory contains a model - config_file = item / "config.json" - adapter_config = item / "adapter_config.json" - - if config_file.exists() or adapter_config.exists(): - # This is a valid checkpoint - display_name = item.name - checkpoint_path = str(item) - checkpoints.append((display_name, checkpoint_path)) - logger.debug(f"Found checkpoint: {display_name}") - - # Sort by modification time (newest first) - checkpoints.sort(key=lambda x: Path(x[1]).stat().st_mtime, reverse=True) - - logger.info(f"Found {len(checkpoints)} checkpoints in {outputs_dir}") - return checkpoints - - except Exception as e: - logger.error(f"Error scanning checkpoints: {e}") - return [] - - def load_checkpoint(self, - checkpoint_path: str, - max_seq_length: int = 2048, - load_in_4bit: bool = True) -> Tuple[bool, str]: - """ - Load a checkpoint for export. - - Returns: - Tuple of (success: bool, message: str) - """ - try: - logger.info(f"Loading checkpoint: {checkpoint_path}") - - # First, cleanup existing models - self.cleanup_memory() - - # Detect if vision model - checkpoint_path_obj = Path(checkpoint_path) - - # Check if it's a LoRA adapter - adapter_config = checkpoint_path_obj / "adapter_config.json" - if adapter_config.exists(): - # It's a LoRA - get base model to check vision - base_model = get_base_model_from_lora(checkpoint_path) - if base_model: - self.is_vision = is_vision_model(base_model) - else: - return False, "Could not determine base model for adapter" - else: - # Check the model itself - self.is_vision = is_vision_model(checkpoint_path) - - # Load model based on type - if self.is_vision: - logger.info("Loading as vision model...") - model, processor = FastVisionModel.from_pretrained( - model_name=checkpoint_path, - max_seq_length=max_seq_length, - dtype=None, - load_in_4bit=load_in_4bit, - ) - tokenizer = processor # For vision models, processor acts as tokenizer - else: - logger.info("Loading as text model...") - model, tokenizer = FastLanguageModel.from_pretrained( - model_name=checkpoint_path, - max_seq_length=max_seq_length, - dtype=None, - load_in_4bit=load_in_4bit, - ) - - # Check if PEFT model - self.is_peft = isinstance(model, (PeftModel, PeftModelForCausalLM)) - - # Store loaded model - self.current_model = model - self.current_tokenizer = tokenizer - self.current_checkpoint = checkpoint_path - - model_type = "Vision" if self.is_vision else "Text" - peft_info = " (PEFT Adapter)" if self.is_peft else " (Merged Model)" - - logger.info(f"Successfully loaded {model_type} model{peft_info}") - return True, f"Loaded {model_type} model{peft_info} successfully" - - except Exception as e: - logger.error(f"Error loading checkpoint: {e}") - import traceback - logger.error(traceback.format_exc()) - return False, f"Failed to load checkpoint: {str(e)}" - - def export_merged_model(self, - save_directory: str, - format_type: str = "16-bit (FP16)", - push_to_hub: bool = False, - repo_id: Optional[str] = None, - hf_token: Optional[str] = None, - private: bool = False) -> Tuple[bool, str]: - """ - Export merged model (for PEFT models). - - Args: - save_directory: Local directory to save model - format_type: "16-bit (FP16)" or "4-bit (FP4)" - push_to_hub: Whether to push to Hugging Face Hub - repo_id: Hub repository ID (username/model-name) - hf_token: Hugging Face token - private: Whether to make the repo private - - Returns: - Tuple of (success: bool, message: str) - """ - if not self.current_model or not self.current_tokenizer: - return False, "No model loaded. Please select a checkpoint first." - - if not self.is_peft: - return False, "This is not a PEFT model. Use 'Export Base Model' instead." - - try: - # Determine save method - if format_type == "4-bit (FP4)": - save_method = "merged_4bit_forced" - else: # 16-bit (FP16) - save_method = "merged_16bit" - - # Save locally if requested - if save_directory: - logger.info(f"Saving merged model locally to: {save_directory}") - os.makedirs(save_directory, exist_ok=True) - - self.current_model.save_pretrained_merged( - save_directory, - self.current_tokenizer, - save_method=save_method - ) - logger.info(f"Model saved successfully to {save_directory}") - - # Push to hub if requested - if push_to_hub: - if not repo_id or not hf_token: - return False, "Repository ID and Hugging Face token required for Hub upload" - - logger.info(f"Pushing merged model to Hub: {repo_id}") - - self.current_model.push_to_hub_merged( - repo_id, - self.current_tokenizer, - save_method=save_method, - token=hf_token, - private=private - ) - logger.info(f"Model pushed successfully to {repo_id}") - - return True, "Model exported successfully" - - except Exception as e: - logger.error(f"Error exporting merged model: {e}") - import traceback - logger.error(traceback.format_exc()) - return False, f"Export failed: {str(e)}" - - def export_base_model(self, - save_directory: str, - push_to_hub: bool = False, - repo_id: Optional[str] = None, - hf_token: Optional[str] = None, - private: bool = False) -> Tuple[bool, str]: - """ - Export base model (for non-PEFT models). - - Returns: - Tuple of (success: bool, message: str) - """ - if not self.current_model or not self.current_tokenizer: - return False, "No model loaded. Please select a checkpoint first." - - if self.is_peft: - return False, "This is a PEFT model. Use 'Merged Model' export type instead." - - try: - # Save locally if requested - if save_directory: - logger.info(f"Saving base model locally to: {save_directory}") - os.makedirs(save_directory, exist_ok=True) - - self.current_model.save_pretrained(save_directory) - self.current_tokenizer.save_pretrained(save_directory) - logger.info(f"Model saved successfully to {save_directory}") - - # Push to hub if requested - if push_to_hub: - if not repo_id or not hf_token: - return False, "Repository ID and Hugging Face token required for Hub upload" - - logger.info(f"Pushing base model to Hub: {repo_id}") - - # Get base model name - base_model = self.current_model.config._name_or_path - - # Create repo - hf_api = HfApi(token=hf_token) - repo_id = PushToHubMixin._create_repo( - PushToHubMixin, - repo_id=repo_id, - private=private, - token=hf_token, - ) - username = repo_id.split("/")[0] - - # Create and push model card - content = MODEL_CARD.format( - username=username, - base_model=base_model, - model_type=self.current_model.config.model_type, - method="", - extra="unsloth", - ) - card = ModelCard(content) - card.push_to_hub(repo_id, token=hf_token, commit_message="Unsloth Model Card") - - # Upload model files - if save_directory: - hf_api.upload_folder( - folder_path=save_directory, - repo_id=repo_id, - repo_type="model" - ) - logger.info(f"Model pushed successfully to {repo_id}") - else: - return False, "Local save directory required for Hub upload" - - return True, "Model exported successfully" - - except Exception as e: - logger.error(f"Error exporting base model: {e}") - import traceback - logger.error(traceback.format_exc()) - return False, f"Export failed: {str(e)}" - - - def export_gguf(self, - save_directory: str, - quantization_method: str = "Q4_K_M", - push_to_hub: bool = False, - repo_id: Optional[str] = None, - hf_token: Optional[str] = None) -> Tuple[bool, str]: - """ - Export model in GGUF format. - - Args: - save_directory: Local directory to save model - quantization_method: GGUF quantization method (e.g., "Q4_K_M") - push_to_hub: Whether to push to Hugging Face Hub - repo_id: Hub repository ID - hf_token: Hugging Face token - - Returns: - Tuple of (success: bool, message: str) - """ - if not self.current_model or not self.current_tokenizer: - return False, "No model loaded. Please select a checkpoint first." - - try: - # Convert quantization method to lowercase for unsloth - quant_method = quantization_method.lower() - - # Save locally if requested - if save_directory: - logger.info(f"Saving GGUF model locally to: {save_directory}") - - # Create the directory if it doesn't exist - os.makedirs(save_directory, exist_ok=True) - - # Get the base filename for the GGUF file - import shutil - original_dir = os.getcwd() - - try: - # Change to target directory - os.chdir(save_directory) - logger.info(f"Changed directory to: {save_directory}") - - # Now save (will save in current directory) - self.current_model.save_pretrained_gguf( - "model", # Base filename - self.current_tokenizer, - quantization_method=quant_method - ) - - logger.info(f"GGUF model saved successfully in {save_directory}") - - # Check if llama.cpp directory was created here - llama_cpp_in_target = os.path.join(save_directory, "llama.cpp") - llama_cpp_in_original = os.path.join(original_dir, "llama.cpp") - - if os.path.exists(llama_cpp_in_target): - logger.info(f"Found llama.cpp directory in {save_directory}") - - # Remove llama.cpp from original directory if it exists - if os.path.exists(llama_cpp_in_original): - logger.info(f"Removing existing llama.cpp in {original_dir}") - shutil.rmtree(llama_cpp_in_original) - - # Move llama.cpp back to original directory - logger.info(f"Moving llama.cpp to {original_dir}") - shutil.move(llama_cpp_in_target, llama_cpp_in_original) - logger.info(f"Successfully moved llama.cpp back to original directory") - - finally: - # Always change back to original directory - os.chdir(original_dir) - logger.info(f"Changed back to original directory: {original_dir}") - - # Push to hub if requested - if push_to_hub: - if not repo_id or not hf_token: - return False, "Repository ID and Hugging Face token required for Hub upload" - - logger.info(f"Pushing GGUF model to Hub: {repo_id}") - - self.current_model.push_to_hub_gguf( - repo_id, - self.current_tokenizer, - quantization_method=quant_method, - token=hf_token - ) - logger.info(f"GGUF model pushed successfully to {repo_id}") - - return True, f"GGUF model exported successfully ({quantization_method})" - - except Exception as e: - logger.error(f"Error exporting GGUF model: {e}") - import traceback - logger.error(traceback.format_exc()) - return False, f"GGUF export failed: {str(e)}" - - def export_lora_adapter(self, - save_directory: str, - push_to_hub: bool = False, - repo_id: Optional[str] = None, - hf_token: Optional[str] = None, - private: bool = False) -> Tuple[bool, str]: - """ - Export LoRA adapter only (not merged). - - Returns: - Tuple of (success: bool, message: str) - """ - if not self.current_model or not self.current_tokenizer: - return False, "No model loaded. Please select a checkpoint first." - - if not self.is_peft: - return False, "This is not a PEFT model. No adapter to export." - - try: - # Save locally if requested - if save_directory: - logger.info(f"Saving LoRA adapter locally to: {save_directory}") - os.makedirs(save_directory, exist_ok=True) - - self.current_model.save_pretrained(save_directory) - self.current_tokenizer.save_pretrained(save_directory) - logger.info(f"Adapter saved successfully to {save_directory}") - - # Push to hub if requested - if push_to_hub: - if not repo_id or not hf_token: - return False, "Repository ID and Hugging Face token required for Hub upload" - - logger.info(f"Pushing LoRA adapter to Hub: {repo_id}") - - self.current_model.push_to_hub( - repo_id, - token=hf_token, - private=private - ) - self.current_tokenizer.push_to_hub( - repo_id, - token=hf_token, - private=private - ) - logger.info(f"Adapter pushed successfully to {repo_id}") - - return True, "LoRA adapter exported successfully" - - except Exception as e: - logger.error(f"Error exporting LoRA adapter: {e}") - import traceback - logger.error(traceback.format_exc()) - return False, f"Adapter export failed: {str(e)}" - - -# Global export backend instance -_export_backend = None - -def get_export_backend() -> ExportBackend: - """Get or create the global export backend instance""" - global _export_backend - if _export_backend is None: - _export_backend = ExportBackend() - return _export_backend diff --git a/studio/backend/backend/inference.py b/studio/backend/backend/inference.py deleted file mode 100644 index 117487cb79..0000000000 --- a/studio/backend/backend/inference.py +++ /dev/null @@ -1,1212 +0,0 @@ -""" -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, Generator, Tuple -from .model_config import ModelConfig, get_base_model_from_lora -from .path_utils import is_model_cached -from utils.utils import format_error_message, 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 = "cuda" if torch.cuda.is_available() else "cpu" - - # 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 - - # Use garbage collection and clear CUDA cache to release memory - import gc - import torch - gc.collect() - if torch.cuda.is_available(): - torch.cuda.empty_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. This returns the base model object. - # This step is only necessary if the model is currently a PeftModel instance. - if isinstance(model, (PeftModel, PeftModelForCausalLM)): - logger.info("Model is a PeftModel. Unloading adapters...") - unwrapped_base_model = model.unload() - self.models[base_model_name]["model"] = unwrapped_base_model - model = unwrapped_base_model # Continue with the unwrapped 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) - - logger.info("Model has been successfully reverted to a 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 .model_config 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 .model_config 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 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. - - 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 - """ - 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") - - 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 - - # Step 1: Apply get_chat_template if model is in mapper - try: - from utils.datasets.dataset_utils 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.""" - 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']*>', '', 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("cuda") - else: - # Text-only for vision model - formatted_prompt = self.format_chat_prompt(messages, system_prompt) - inputs = processor.tokenizer(formatted_prompt, return_tensors="pt").to("cuda") - - # Generate with streaming - captured_output = StringIO() - original_stdout = sys.stdout - - try: - sys.stdout = captured_output - - text_streamer = TextStreamer(processor.tokenizer, skip_prompt=True) - model.generate( - **inputs, - streamer=text_streamer, - max_new_tokens=max_new_tokens, - use_cache=True, - temperature=temperature, - top_p=top_p, - top_k=top_k - ) - - sys.stdout = original_stdout - generated_text = captured_output.getvalue() - cleaned = self._clean_generated_text(generated_text) - yield cleaned - - except Exception as e: - sys.stdout = original_stdout - 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", "") - 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, show_warning: bool = True) -> bool: - """ - Check if current model supports vision and optionally show warning if image uploaded to text-only model - - Args: - show_warning: Whether to show Gradio warning if vision not supported - - 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: - is_vision = self.models[current_model].get("is_vision", False) - if not is_vision and show_warning: - import gradio as gr - model_short = current_model.split('/')[-1] if '/' in current_model else current_model - gr.Warning(f"Image uploaded, but {model_short} is a text-only model. Please select a vision model to analyze images.") - return is_vision - 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) - - import torch - if torch.cuda.is_available(): - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - logger.debug("Cleared CUDA cache and IPC resources") - - 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'|', '', 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.dataset_utils 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 add_local_model_to_dropdown(self, model_path: str): - """Add successfully loaded local model to dropdown storage""" - try: - from pathlib import Path - - path_obj = Path(model_path) - display_name = f"{path_obj.name}" - - # Check if already exists - for existing_display, existing_path in self.loaded_local_models: - if existing_path == model_path: - logger.debug(f"Local model already in dropdown: {model_path}") - return - - # Add to beginning of list - self.loaded_local_models.insert(0, (display_name, model_path)) - logger.info(f"Added local model to dropdown: {display_name} -> {model_path}") - - # Keep only last 5 - if len(self.loaded_local_models) > 5: - self.loaded_local_models.pop() - - except Exception as e: - logger.error(f"Error adding local model to dropdown: {e}") - - def get_model_dropdown_choices(self, models: list = None) -> list: - """Get model dropdown choices with status indicators""" - if models is None: - models = self.default_models - - try: - active_model = self.active_model_name - loading_model = self.get_loading_model() - - choices = [] - - # Add local models first - for local_display, local_path in self.loaded_local_models: - if local_path == active_model: - choices.append((f"{local_display} (Active)", local_path)) - else: - choices.append((local_display, local_path)) - - # Add default models - for model in models: - short_name = model.split('/')[-1] if '/' in model else model - - if model == active_model: - display_name = f"{short_name} (Active)" - elif model == loading_model: - display_name = f"{short_name} (Loading...)" - elif model in self.models and self.models[model].get("model"): - # Model is loaded in memory - display_name = f"{short_name} (Ready)" - # elif model in self.models: - # display_name = f"{short_name} (Ready)" - elif is_model_cached(model): - # Model is downloaded but not loaded - display_name = f"{short_name} (Cached)" - else: - display_name = f"↓ {short_name}" # Not downloaded - - choices.append((display_name, model)) - - return choices - - except Exception as e: - logger.error(f"Error getting model choices: {e}") - return [(model.split('/')[-1], model) for model in models] - - - def update_model_dropdown(self, models: list = None): - """Update model dropdown with current status""" - try: - import gradio as gr - - choices = self.get_model_dropdown_choices(models) - active_model = self.active_model_name - - # Set value to active model if exists - value = active_model if active_model else (choices[0][1] if choices else None) - - return gr.update(choices=choices, value=value) - - except Exception as e: - logger.error(f"Error updating model dropdown: {e}") - import gradio as gr - return gr.update() - - 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 diff --git a/studio/backend/backend/model_config.py b/studio/backend/backend/model_config.py deleted file mode 100644 index 85bdd2cf42..0000000000 --- a/studio/backend/backend/model_config.py +++ /dev/null @@ -1,704 +0,0 @@ -""" -Model and LoRA configuration handling -""" -from transformers import AutoConfig -from dataclasses import dataclass -from typing import Optional, Dict, Any -from .path_utils import normalize_path, is_local_path, is_model_cached -from utils.utils import without_hf_auth -import logging -from pathlib import Path -from typing import List, Tuple -import json -import yaml - - -logger = logging.getLogger(__name__) - -# Model name mapping: maps all equivalent model names to their canonical YAML config file -# Format: "canonical_model_name.yaml": [list of all equivalent model names] -# Based on the model mapper provided - canonical filename is based on the first model name in the mapper -MODEL_NAME_MAPPING = { - "unsloth_answerdotai_ModernBERT-large.yaml": [ - "answerdotai/ModernBERT-large", - ], - "unsloth_Qwen2.5-Coder-7B-Instruct-bnb-4bit.yaml": [ - "unsloth/Qwen2.5-Coder-7B-Instruct-bnb-4bit", - "unsloth/Qwen2.5-Coder-7B-Instruct", - "Qwen/Qwen2.5-Coder-7B-Instruct", - ], - "unsloth_codegemma-7b-bnb-4bit.yaml": [ - "unsloth/codegemma-7b-bnb-4bit", - "unsloth/codegemma-7b", - "google/codegemma-7b", - ], - "unsloth_ERNIE-4.5-21B-A3B-PT.yaml": [ - "unsloth/ERNIE-4.5-21B-A3B-PT", - ], - "unsloth_ERNIE-4.5-VL-28B-A3B-PT.yaml": [ - "unsloth/ERNIE-4.5-VL-28B-A3B-PT", - ], - - "tiiuae_Falcon-H1-0.5B-Instruct.yaml": [ - "tiiuae/Falcon-H1-0.5B-Instruct", - "unsloth/Falcon-H1-0.5B-Instruct", - ], - "unsloth_functiongemma-270m-it.yaml": [ - "unsloth/functiongemma-270m-it-unsloth-bnb-4bit", - "google/functiongemma-270m-it", - "unsloth/functiongemma-270m-it-unsloth-bnb-4bit", - ], - "unsloth_gemma-2-2b.yaml": [ - "unsloth/gemma-2-2b-bnb-4bit", - "google/gemma-2-2b", - ], - "unsloth_gemma-2-27b-bnb-4bit.yaml": [ - "unsloth/gemma-2-9b-bnb-4bit", - "unsloth/gemma-2-9b", - "google/gemma-2-9b", - "unsloth/gemma-2-27b", - "google/gemma-2-27b", - ], - "unsloth_gemma-3-4b-pt.yaml": [ - "unsloth/gemma-3-4b-pt-unsloth-bnb-4bit", - "google/gemma-3-4b-pt", - "unsloth/gemma-3-4b-pt-bnb-4bit", - ], - "unsloth_gemma-3-4b-it.yaml": [ - "unsloth/gemma-3-4b-it-unsloth-bnb-4bit", - "google/gemma-3-4b-it", - "unsloth/gemma-3-4b-it-bnb-4bit", - ], - "unsloth_gemma-3-27b-it.yaml": [ - "unsloth/gemma-3-27b-it-unsloth-bnb-4bit", - "google/gemma-3-27b-it", - "unsloth/gemma-3-27b-it-bnb-4bit", - ], - "unsloth_gemma-3-270m-it.yaml": [ - "unsloth/gemma-3-270m-it-unsloth-bnb-4bit", - "google/gemma-3-270m-it", - "unsloth/gemma-3-270m-it-bnb-4bit", - ], - "unsloth_gemma-3n-E4B-it.yaml": [ - "unsloth/gemma-3n-E4B-it-unsloth-bnb-4bit", - "google/gemma-3n-E4B-it", - "unsloth/gemma-3n-E4B-it-unsloth-bnb-4bit", - ], - "unsloth_gemma-3n-E4B.yaml": [ - "unsloth/gemma-3n-E4B-unsloth-bnb-4bit", - "google/gemma-3n-E4B", - ], - "unsloth_gpt-oss-20b.yaml": [ - "openai/gpt-oss-20b", - "unsloth/gpt-oss-20b-unsloth-bnb-4bit", - "unsloth/gpt-oss-20b-BF16", - ], - - "unsloth_gpt-oss-120b.yaml": [ - "openai/gpt-oss-120b", - "unsloth/gpt-oss-120b-unsloth-bnb-4bit", - ], - "unsloth_granite-4.0-350m-unsloth-bnb-4bit.yaml": [ - "unsloth/granite-4.0-350m", - "ibm-granite/granite-4.0-350m", - "unsloth/granite-4.0-350m-bnb-4bit", - ], - "unsloth_granite-4.0-h-micro.yaml": [ - "ibm-granite/granite-4.0-h-micro", - "unsloth/granite-4.0-h-micro-bnb-4bit", - "unsloth/granite-4.0-h-micro-unsloth-bnb-4bit", - ], - "unsloth_LFM2-1.2B.yaml": [ - "unsloth/LFM2-1.2B", - ], - "unsloth_llama-3-8b-bnb-4bit.yaml": [ - "unsloth/llama-3-8b", - "meta-llama/Meta-Llama-3-8B", - ], - "unsloth_llama-3-8b-Instruct-bnb-4bit.yaml": [ - "unsloth/llama-3-8b-Instruct", - "meta-llama/Meta-Llama-3-8B-Instruct", - ], - "unsloth_Meta-Llama-3.1-70B-bnb-4bit.yaml": [ - "unsloth/Meta-Llama-3.1-8B-bnb-4bit", - "unsloth/Meta-Llama-3.1-8B-unsloth-bnb-4bit", - "meta-llama/Meta-Llama-3.1-8B", - "unsloth/Meta-Llama-3.1-70B-bnb-4bit", - "unsloth/Meta-Llama-3.1-8B", - "unsloth/Meta-Llama-3.1-70B", - "meta-llama/Meta-Llama-3.1-70B", - "unsloth/Meta-Llama-3.1-405B-bnb-4bit", - "meta-llama/Meta-Llama-3.1-405B", - ], - - "unsloth_Meta-Llama-3.1-8B-Instruct-bnb-4bit.yaml": [ - "unsloth/Meta-Llama-3.1-8B-Instruct-unsloth-bnb-4bit", - "unsloth/Meta-Llama-3.1-8B-Instruct-bnb-4bit", - "meta-llama/Meta-Llama-3.1-8B-Instruct", - "unsloth/Meta-Llama-3.1-8B-Instruct", - "RedHatAI/Llama-3.1-8B-Instruct-FP8", - "unsloth/Llama-3.1-8B-Instruct-FP8-Block", - "unsloth/Llama-3.1-8B-Instruct-FP8-Dynamic", - ], - "unsloth_Llama-3.2-3B-Instruct.yaml": [ - "unsloth/Llama-3.2-3B-Instruct-unsloth-bnb-4bit", - "meta-llama/Llama-3.2-3B-Instruct", - "unsloth/Llama-3.2-3B-Instruct-bnb-4bit", - "RedHatAI/Llama-3.2-3B-Instruct-FP8", - "unsloth/Llama-3.2-3B-Instruct-FP8-Block", - "unsloth/Llama-3.2-3B-Instruct-FP8-Dynamic", - ], - "unsloth_Llama-3.2-1B-Instruct.yaml": [ - "unsloth/Llama-3.2-1B-Instruct-unsloth-bnb-4bit", - "meta-llama/Llama-3.2-1B-Instruct", - "unsloth/Llama-3.2-1B-Instruct-bnb-4bit", - "RedHatAI/Llama-3.2-1B-Instruct-FP8", - "unsloth/Llama-3.2-1B-Instruct-FP8-Block", - "unsloth/Llama-3.2-1B-Instruct-FP8-Dynamic", - ], - "unsloth_Llama-3.2-11B-Vision-Instruct.yaml": [ - "unsloth/Llama-3.2-11B-Vision-Instruct-unsloth-bnb-4bit", - "meta-llama/Llama-3.2-11B-Vision-Instruct", - "unsloth/Llama-3.2-11B-Vision-Instruct-bnb-4bit", - ], - "unsloth_Llama-3.3-70B-Instruct.yaml": [ - "unsloth/Llama-3.3-70B-Instruct-unsloth-bnb-4bit", - "meta-llama/Llama-3.3-70B-Instruct", - "unsloth/Llama-3.3-70B-Instruct-bnb-4bit", - "RedHatAI/Llama-3.3-70B-Instruct-FP8", - "unsloth/Llama-3.3-70B-Instruct-FP8-Block", - "unsloth/Llama-3.3-70B-Instruct-FP8-Dynamic", - ], - "unsloth_Llasa-3B.yaml": [ - "HKUSTAudio/Llasa-1B", - "unsloth/Llasa-3B", - ], - "unsloth_Magistral-Small-2509-unsloth-bnb-4bit.yaml": [ - "unsloth/Magistral-Small-2509", - "mistralai/Magistral-Small-2509", - "unsloth/Magistral-Small-2509-bnb-4bit", - ], - "unsloth_Ministral-3-3B-Instruct-2512.yaml": [ - "unsloth/Ministral-3-3B-Instruct-2512", - ], - "unsloth_mistral-7b-v0.3-bnb-4bit.yaml": [ - "unsloth/mistral-7b-v0.3-bnb-4bit" - "unsloth/mistral-7b-v0.3", - "mistralai/Mistral-7B-v0.3", - ], - "unsloth_Mistral-Nemo-Base-2407-bnb-4bit.yaml": [ - "unsloth/Mistral-Nemo-Base-2407-bnb-4bit", - "unsloth/Mistral-Nemo-Base-2407", - "mistralai/Mistral-Nemo-Base-2407", - "unsloth/Mistral-Nemo-Instruct-2407-bnb-4bit", - "unsloth/Mistral-Nemo-Instruct-2407", - "mistralai/Mistral-Nemo-Instruct-2407", - - ], - "unsloth_Mistral-Small-Instruct-2409.yaml": [ - "unsloth/Mistral-Small-Instruct-2409-bnb-4bit", - "mistralai/Mistral-Small-Instruct-2409", - ], - "unsloth_mistral-7b-instruct-v0.3-bnb-4bit.yaml": [ - "unsloth/mistral-7b-instruct-v0.3-bnb-4bit", - "unsloth/mistral-7b-instruct-v0.3", - "mistralai/Mistral-7B-Instruct-v0.3", - ], - "unsloth_Qwen2.5-1.5B-Instruct.yaml": [ - "unsloth/Qwen2.5-1.5B-Instruct-unsloth-bnb-4bit", - "Qwen/Qwen2.5-1.5B-Instruct", - "unsloth/Qwen2.5-1.5B-Instruct-bnb-4bit", - ], - "unsloth_Nemotron-3-Nano-30B-A3B.yaml": [ - "unsloth/Nemotron-3-Nano-30B-A3B", - ], - "unsloth_orpheus-3b-0.1-ft.yaml": [ - "unsloth/orpheus-3b-0.1-ft-unsloth-bnb-4bit", - "canopylabs/orpheus-3b-0.1-ft", - "unsloth/orpheus-3b-0.1-ft-bnb-4bit", - ], - "OuteAI_Llama-OuteTTS-1.0-1B.yaml": [ - "OuteAI/Llama-OuteTTS-1.0-1B", - ], - "unsloth_PaddleOCR-VL.yaml": [ - "unsloth/PaddleOCR-VL", - ], - "unsloth_Phi-3-medium-4k-instruct.yaml": [ - "unsloth/Phi-3-medium-4k-instruct-bnb-4bit", - "microsoft/Phi-3-medium-4k-instruct", - ], - "unsloth_Phi-3.5-mini-instruct.yaml": [ - "unsloth/Phi-3.5-mini-instruct-bnb-4bit", - "microsoft/Phi-3.5-mini-instruct", - ], - "unsloth_Phi-4.yaml": [ - "unsloth/phi-4-unsloth-bnb-4bit", - "microsoft/phi-4", - "unsloth/phi-4-bnb-4bit", - ], - "unsloth_Pixtral-12B-2409.yaml": [ - "unsloth/Pixtral-12B-2409-unsloth-bnb-4bit", - "mistralai/Pixtral-12B-2409", - "unsloth/Pixtral-12B-2409-bnb-4bit", - ], - "unsloth_Qwen2-7B.yaml": [ - "unsloth/Qwen2-7B-bnb-4bit", - "Qwen/Qwen2-7B", - ], - "unsloth_Qwen2-VL-7B-Instruct.yaml": [ - "unsloth/Qwen2-VL-7B-Instruct-unsloth-bnb-4bit", - "Qwen/Qwen2-VL-7B-Instruct", - "unsloth/Qwen2-VL-7B-Instruct-bnb-4bit", - ], - "unsloth_Qwen2.5-7B.yaml": [ - "unsloth/Qwen2.5-7B-unsloth-bnb-4bit", - "Qwen/Qwen2.5-7B", - "unsloth/Qwen2.5-7B-bnb-4bit", - ], - "unsloth_Qwen2.5-Coder-1.5B-Instruct.yaml": [ - "unsloth/Qwen2.5-Coder-1.5B-Instruct-bnb-4bit", - "Qwen/Qwen2.5-Coder-1.5B-Instruct", - ], - "unsloth_Qwen2.5-Coder-14B-Instruct.yaml": [ - "unsloth/Qwen2.5-Coder-14B-Instruct-bnb-4bit", - "Qwen/Qwen2.5-Coder-14B-Instruct", - ], - "unsloth_Qwen2.5-VL-7B-Instruct-bnb-4bit.yaml": [ - "unsloth/Qwen2.5-VL-7B-Instruct", - "Qwen/Qwen2.5-VL-7B-Instruct", - "unsloth/Qwen2.5-VL-7B-Instruct-unsloth-bnb-4bit", - ], - "unsloth_Qwen3-0.6B.yaml": [ - "unsloth/Qwen3-0.6B-unsloth-bnb-4bit", - "Qwen/Qwen3-0.6B", - "unsloth/Qwen3-0.6B-bnb-4bit", - "Qwen/Qwen3-0.6B-FP8", - "unsloth/Qwen3-0.6B-FP8", - ], - "unsloth_Qwen3-4B-Instruct-2507.yaml": [ - "unsloth/Qwen3-4B-Instruct-2507-unsloth-bnb-4bit", - "Qwen/Qwen3-4B-Instruct-2507", - "unsloth/Qwen3-4B-Instruct-2507-bnb-4bit", - "Qwen/Qwen3-4B-Instruct-2507-FP8", - "unsloth/Qwen3-4B-Instruct-2507-FP8", - ], - "unsloth_Qwen3-4B-Thinking-2507.yaml": [ - "unsloth/Qwen3-4B-Thinking-2507-unsloth-bnb-4bit", - "Qwen/Qwen3-4B-Thinking-2507", - "unsloth/Qwen3-4B-Thinking-2507-bnb-4bit", - "Qwen/Qwen3-4B-Thinking-2507-FP8", - "unsloth/Qwen3-4B-Thinking-2507-FP8", - ], - "unsloth_Qwen3-14B-Base-unsloth-bnb-4bit.yaml": [ - "unsloth/Qwen3-14B-Base", - "Qwen/Qwen3-14B-Base", - "unsloth/Qwen3-14B-Base-bnb-4bit", - ], - "unsloth_Qwen3-14B.yaml": [ - "unsloth/Qwen3-14B-unsloth-bnb-4bit", - "Qwen/Qwen3-14B", - "unsloth/Qwen3-14B-bnb-4bit", - "Qwen/Qwen3-14B-FP8", - "unsloth/Qwen3-14B-FP8", - ], - "unsloth_Qwen3-32B.yaml": [ - "unsloth/Qwen3-32B-unsloth-bnb-4bit", - "Qwen/Qwen3-32B", - "unsloth/Qwen3-32B-bnb-4bit", - "Qwen/Qwen3-32B-FP8", - "unsloth/Qwen3-32B-FP8", - ], - "unsloth_Qwen3-VL-8B-Instruct-unsloth-bnb-4bit.yaml": [ - "Qwen/Qwen3-VL-8B-Instruct-FP8", - "unsloth/Qwen3-VL-8B-Instruct-FP8", - "unsloth/Qwen3-VL-8B-Instruct", - "Qwen/Qwen3-VL-8B-Instruct", - "unsloth/Qwen3-VL-8B-Instruct-bnb-4bit", - ], - "sesame_csm-1b.yaml": [ - "sesame/csm-1b", - ], - "Spark-TTS-0.5B_LLM.yaml": [ - "Spark-TTS-0.5B/LLM", - ], - "unsloth_tinyllama-bnb-4bit.yaml": [ - "unsloth/tinyllama", - "TinyLlama/TinyLlama-1.1B-intermediate-step-1431k-3T", - ], - "unsloth_whisper-large-v3.yaml": [ - "unsloth/whisper-large-v3", - "openai/whisper-large-v3", - ], -} - -# Reverse mapping for quick lookup: model_name -> canonical_filename -_REVERSE_MODEL_MAPPING = {} -for canonical_file, model_names in MODEL_NAME_MAPPING.items(): - for model_name in model_names: - _REVERSE_MODEL_MAPPING[model_name] = canonical_file - -def load_model_config(model_name: str, use_auth: bool = False, token: Optional[str] = None): - """ - Load model config with optional authentication control. - """ - - if token: - # Explicit token provided - use it - return AutoConfig.from_pretrained( - model_name, - trust_remote_code=True, - token=token - ) - - if not use_auth: - # Load without any authentication (for public model checks) - with without_hf_auth(): - return AutoConfig.from_pretrained( - model_name, - trust_remote_code=True, - token=None - ) - - # Use default authentication (cached tokens) - return AutoConfig.from_pretrained( - model_name, - trust_remote_code=True - ) -pass - - -def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool: - """ - Detect vision models by checking architecture in config. - Works for fine-tuned models since they inherit the base architecture. - - Args: - model_name: Model identifier (HF repo or local path) - hf_token: Optional HF token for accessing gated/private models - """ - try: - config = load_model_config(model_name, token=hf_token) - - # Check vision arch - if hasattr(config, 'architectures'): - is_vlm = any( - x.endswith(("ForConditionalGeneration", "ForVisionText2Text")) - for x in config.architectures - ) - if is_vlm: - logger.info(f"Model {model_name} detected as vision model: architecture {config.architectures}") - return True - - # Quick check for vision config as backup - if hasattr(config, 'vision_config'): - logger.info(f"Model {model_name} detected as vision model: has vision_config") - return True - - return False - - except Exception as e: - logger.warning(f"Could not determine if {model_name} is vision model: {e}") - return False -pass - - -def scan_trained_loras(outputs_dir: str = "./outputs") -> List[Tuple[str, str]]: - """ - Scan outputs folder for trained LoRA adapters. - - Returns: - List of tuples: [(display_name, adapter_path), ...] - - Example: - [ - ("unsloth_Meta-Llama-3.1_...", "./outputs/unsloth_Meta-Llama-3.1_.../"), - ("my_finetuned_model", "./outputs/my_finetuned_model/"), - ] - """ - trained_loras = [] - outputs_path = Path(outputs_dir) - - if not outputs_path.exists(): - logger.warning(f"Outputs directory not found: {outputs_dir}") - return trained_loras - - try: - for item in outputs_path.iterdir(): - if item.is_dir(): - # Check if this directory contains a LoRA adapter - adapter_config = item / "adapter_config.json" - adapter_model = item / "adapter_model.safetensors" - - if adapter_config.exists() or adapter_model.exists(): - display_name = item.name - adapter_path = str(item) - trained_loras.append((display_name, adapter_path)) - logger.debug(f"Found trained LoRA: {display_name}") - - # Sort by modification time (newest first) - trained_loras.sort(key=lambda x: Path(x[1]).stat().st_mtime, reverse=True) - - logger.info(f"Found {len(trained_loras)} trained LoRA adapters in {outputs_dir}") - return trained_loras - - except Exception as e: - logger.error(f"Error scanning outputs folder: {e}") - return [] - -def get_base_model_from_lora(lora_path: str) -> Optional[str]: - """ - Read the base model name from a LoRA adapter's config. - - Args: - lora_path: Path to the LoRA adapter directory - - Returns: - Base model identifier (e.g., "unsloth/Meta-Llama-3.1-8B-Instruct-bnb-4bit") - or None if not found - - Example: - >>> get_base_model_from_lora("./outputs/unsloth_Meta-Llama-3.1_.../") - "unsloth/Meta-Llama-3.1-8B-Instruct-bnb-4bit" - """ - try: - lora_path_obj = Path(lora_path) - - # Try adapter_config.json first - adapter_config_path = lora_path_obj / "adapter_config.json" - if adapter_config_path.exists(): - with open(adapter_config_path, 'r') as f: - config = json.load(f) - base_model = config.get("base_model_name_or_path") - if base_model: - logger.info(f"Detected base model from adapter_config.json: {base_model}") - return base_model - - # Fallback: try training_args.bin (requires torch) - training_args_path = lora_path_obj / "training_args.bin" - if training_args_path.exists(): - try: - import torch - training_args = torch.load(training_args_path) - if hasattr(training_args, 'model_name_or_path'): - base_model = training_args.model_name_or_path - logger.info(f"Detected base model from training_args.bin: {base_model}") - return base_model - except Exception as e: - logger.warning(f"Could not load training_args.bin: {e}") - - # Last resort: parse from directory name - # Format: unsloth_Meta-Llama-3.1-8B-Instruct-bnb-4bit_timestamp - dir_name = lora_path_obj.name - if dir_name.startswith("unsloth_"): - # Remove timestamp suffix (usually _1234567890) - parts = dir_name.split("_") - # Reconstruct model name - if len(parts) >= 2: - model_parts = parts[1:-1] # Skip "unsloth" and timestamp - base_model = "unsloth/" + "_".join(model_parts) - logger.info(f"Detected base model from directory name: {base_model}") - return base_model - - logger.warning(f"Could not detect base model for LoRA: {lora_path}") - return None - - except Exception as e: - logger.error(f"Error reading base model from LoRA config: {e}") - return None -pass - -# Status indicators that appear in UI dropdowns -UI_STATUS_INDICATORS = [" (Ready)", " (Loading...)", " (Active)", "↓ "] - -def load_model_defaults(model_name: str) -> Dict[str, Any]: - """ - Load default training parameters for a model from YAML file. - - Args: - model_name: Model identifier (e.g., "unsloth/Meta-Llama-3.1-8B-bnb-4bit") - - Returns: - Dictionary with default parameters from YAML file, or empty dict if not found - - The function looks for a YAML file in configs/model_defaults/ (including subfolders) - based on the model name or its aliases from MODEL_NAME_MAPPING. - If no specific file exists, it falls back to default.yaml. - """ - try: - # Get the script directory to locate configs - script_dir = Path(__file__).parent.parent - defaults_dir = script_dir / "configs" / "model_defaults" - - # First, check if model is in the mapping - if model_name in _REVERSE_MODEL_MAPPING: - canonical_file = _REVERSE_MODEL_MAPPING[model_name] - # Search in subfolders and root - for config_path in defaults_dir.rglob(canonical_file): - if config_path.is_file(): - with open(config_path, 'r', encoding='utf-8') as f: - config = yaml.safe_load(f) or {} - logger.info(f"Loaded model defaults from {config_path} (via mapping)") - return config - - # Try exact model name match (for backward compatibility) - model_filename = model_name.replace("/", "_") + ".yaml" - # Search in subfolders and root - for config_path in defaults_dir.rglob(model_filename): - if config_path.is_file(): - with open(config_path, 'r', encoding='utf-8') as f: - config = yaml.safe_load(f) or {} - logger.info(f"Loaded model defaults from {config_path}") - return config - - # Fall back to default.yaml - default_config_path = defaults_dir / "default.yaml" - if default_config_path.exists(): - with open(default_config_path, 'r', encoding='utf-8') as f: - config = yaml.safe_load(f) or {} - logger.info(f"Loaded default model defaults from {default_config_path}") - return config - - logger.warning(f"No default config found for model {model_name}") - return {} - - except Exception as e: - logger.error(f"Error loading model defaults for {model_name}: {e}") - return {} - -@dataclass -class ModelConfig: - """Configuration for a model to load""" - identifier: str # Clean model identifier (org/name or path) - display_name: str # Original UI display name - path: str # Normalized filesystem path - is_local: bool # Is this a local file vs HF model? - is_cached: bool # Is this already in HF cache? - is_vision: bool # Is this a vision model? - is_lora: bool # Is this a lora adapter? - base_model: Optional[str] = None # Base model (for LoRAs) - - @classmethod - def from_lora_path(cls, lora_path: str, hf_token: Optional[str] = None) -> Optional['ModelConfig']: - """ - Create ModelConfig from a local LoRA adapter path. - - Automatically detects the base model from adapter config. - - Args: - lora_path: Path to LoRA adapter (e.g., "./outputs/unsloth_Meta-Llama-3.1_.../") - hf_token: HF token for vision detection - - Returns: - ModelConfig for the LoRA adapter - """ - try: - lora_path_obj = Path(lora_path) - - if not lora_path_obj.exists(): - logger.error(f"LoRA path does not exist: {lora_path}") - return None - - # Get base model - base_model = get_base_model_from_lora(lora_path) - if not base_model: - logger.error(f"Could not determine base model for LoRA: {lora_path}") - return None - - # Check if base model is vision - is_vision = is_vision_model(base_model, hf_token=hf_token) - - display_name = lora_path_obj.name - identifier = lora_path # Use path as identifier for local LoRAs - - return cls( - identifier=identifier, - display_name=display_name, - path=lora_path, - is_local=True, - is_cached=True, # Local LoRAs are always "cached" - is_vision=is_vision, - is_lora=True, - base_model=base_model, - ) - - except Exception as e: - logger.error(f"Error creating ModelConfig from LoRA path: {e}") - return None - - - @classmethod - def from_ui_selection(cls, - dropdown_value: Optional[str], - search_value: Optional[str], - local_models: list = None, - hf_token: Optional[str] = None, - is_lora: bool = False) -> Optional['ModelConfig']: - """ - Create a universal ModelConfig from UI dropdown/search selections. - Handles base models and LoRA adapters. - """ - selected = None - if search_value and search_value.strip(): - selected = search_value.strip() - elif dropdown_value: - selected = dropdown_value - - if not selected: - return None - - display_name = selected - - # Use the correct 'local_models' parameter to resolve display names - if " (Active)" in selected or " (Ready)" in selected: - clean_display_name = selected.replace(" (Active)", "").replace(" (Ready)", "") - if local_models: - for local_display, local_path in local_models: - if local_display == clean_display_name: - selected = local_path - break - - # Clean all UI status indicators to get the final identifier - identifier = selected - for status in UI_STATUS_INDICATORS: - identifier = identifier.replace(status, "") - identifier = identifier.strip() - - is_local = is_local_path(identifier) - path = normalize_path(identifier) if is_local else identifier - - # Add unsloth/ prefix for shorthand HF models - if not is_local and "/" not in identifier: - identifier = f"unsloth/{identifier}" - path = identifier - - # --- Logic for Base Model and Vision Detection --- - base_model = None - is_vision = False - - if is_lora: - # For a LoRA, we MUST find its base model. - base_model = get_base_model_from_lora(path) - if not base_model: - logger.warning(f"Could not determine base model for LoRA '{path}'. Cannot create config.") - return None # Cannot proceed without a base model - - # A LoRA's vision capability is determined by its base model. - is_vision = is_vision_model(base_model, hf_token=hf_token) - else: - # For a base model, just check its own vision status. - is_vision = is_vision_model(identifier, hf_token=hf_token) - - from .path_utils import is_model_cached - is_cached = is_model_cached(identifier) if not is_local else True - - return cls( - identifier=identifier, - display_name=display_name, - path=path, - is_local=is_local, - is_cached=is_cached, - is_vision=is_vision, - is_lora=is_lora, - base_model=base_model, # This will be None for base models, and populated for LoRAs - ) - pass diff --git a/studio/backend/backend/path_utils.py b/studio/backend/backend/path_utils.py deleted file mode 100644 index 7743952b6b..0000000000 --- a/studio/backend/backend/path_utils.py +++ /dev/null @@ -1,78 +0,0 @@ -""" -Path utilities for model and dataset handling -""" -import os -from pathlib import Path -from typing import Optional -import logging - -logger = logging.getLogger(__name__) - - -def normalize_path(path: str) -> str: - """ - Convert Windows paths to WSL format if needed. - - Examples: - C:\\Users\\... -> /mnt/c/Users/... - /home/user/... -> /home/user/... (unchanged) - """ - if not path: - return path - - # Handle Windows drive letters (C:\\ or c:\\) - if len(path) >= 3 and path[1] == ':' and path[2] in ('\\', '/'): - drive = path[0].lower() - rest = path[3:].replace('\\', '/') - return f'/mnt/{drive}/{rest}' - - # Already Unix-style or relative - return path.replace('\\', '/') -pass - -def is_local_path(path: str) -> bool: - """ - Check if path is a local filesystem path vs HuggingFace model identifier. - - Examples: - True: /home/user/model, C:\\models, ./model, ~/model - False: unsloth/llama-3.1-8b, microsoft/phi-2 - """ - if not path: - return False - - # Obvious HF patterns - if path.count('/') == 1 and not path.startswith(('/', '.', '~')): - return False # Looks like org/model format - - # Filesystem indicators - return ( - path.startswith(('/', '.', '~')) or # Unix absolute/relative - ':' in path or # Windows drive or URL - '\\' in path or # Windows separator - os.path.isabs(path) # System-absolute - ) -pass - -def get_cache_path(model_name: str) -> Optional[Path]: - """Get HuggingFace cache path for a model if it exists.""" - cache_dir = Path.home() / '.cache' / 'huggingface' / 'hub' - model_cache_name = model_name.replace("/", "--") - model_cache_path = cache_dir / f'models--{model_cache_name}' - - return model_cache_path if model_cache_path.exists() else None -pass - -def is_model_cached(model_name: str) -> bool: - """Check if model is downloaded in HuggingFace cache.""" - cache_path = get_cache_path(model_name) - if not cache_path: - return False - - # Check for actual model files - for suffix in ['.safetensors', '.bin', '.json']: - if list(cache_path.rglob(f'*{suffix}')): - return True - - return False -pass diff --git a/studio/backend/backend/trainer.py b/studio/backend/backend/trainer.py deleted file mode 100644 index ce66e8ea6f..0000000000 --- a/studio/backend/backend/trainer.py +++ /dev/null @@ -1,872 +0,0 @@ -""" -Unsloth Training Backend -Integrates Unsloth training capabilities with the Gradio UI -""" -import torch -torch._dynamo.config.recompile_limit = 64 -from unsloth import FastLanguageModel, FastVisionModel, is_bfloat16_supported -from unsloth.chat_templates import get_chat_template - -import os -import json -import threading -import math -import logging -from typing import Optional, Callable -from dataclasses import dataclass -import pandas as pd -from datasets import Dataset, load_dataset - -# Add the parent directory to sys.path to import unsloth modules -#sys.path.append(os.path.join(os.path.dirname(__file__), '..')) -from .model_config import is_vision_model -from utils.datasets.dataset_utils import format_and_template_dataset -from utils.datasets.dataset_utils import MODEL_TO_TEMPLATE_MAPPER, TEMPLATE_TO_RESPONSES_MAPPER -from trl import SFTTrainer, SFTConfig - -# Import Unsloth trainers -#from unsloth_compiled_cache.UnslothSFTTrainer import _UnslothSFTTrainer as SFTTrainer - -logging.basicConfig(level=logging.INFO) -logger = logging.getLogger(__name__) - -@dataclass -class TrainingProgress: - """Training progress tracking""" - epoch: int = 0 - step: int = 0 - total_steps: int = 0 - loss: float = 0.0 - learning_rate: float = 0.0 - is_training: bool = False - is_completed: bool = False - error: Optional[str] = None - status_message: str = "Ready to train" # Current stage message - -class UnslothTrainer: - """ - Unsloth Training Backend for Gradio UI Integration - """ - - def __init__(self): - self.model = None - self.tokenizer = None - self.trainer = None - self.training_thread = None - self.training_progress = TrainingProgress() - self.progress_callbacks = [] - self.is_training = False - self.should_stop = False - - # Model state tracking - self.is_vlm = False - self.model_name = None - - # Thread safety - self._lock = threading.Lock() - - # Store training context for later transfer - self.training_context = { - 'base_model_name': None, - 'output_dir': None, - 'is_lora': True, # Default to LoRA - } - - def add_progress_callback(self, callback: Callable[[TrainingProgress], None]): - """Add callback for training progress updates""" - self.progress_callbacks.append(callback) - - def _update_progress(self, **kwargs): - """Update training progress and notify callbacks""" - with self._lock: - for key, value in kwargs.items(): - if hasattr(self.training_progress, key): - setattr(self.training_progress, key, value) - - # Notify all callbacks - for callback in self.progress_callbacks: - try: - callback(self.training_progress) - except Exception as e: - logger.error(f"Error in progress callback: {e}") - - def load_model(self, - model_name: str, - max_seq_length: int = 2048, - load_in_4bit: bool = True, - hf_token: Optional[str] = None) -> bool: - """Load model for training (supports both text and vision models)""" - try: - print("\nClearing GPU memory before training...") - torch.cuda.empty_cache() - import gc - gc.collect() - - # Detect if this is a vision model first - self.is_vlm = is_vision_model(model_name) - self.model_name = model_name - - logger.info(f"Model type detected: {'Vision' if self.is_vlm else 'Text'}") - - # Reset training state for new run - self._update_progress( - is_training=True, - is_completed=False, - error=None, - step=0, - loss=0.0, - epoch=0 - ) - - # Update UI immediately with loading message - model_display = model_name.split('/')[-1] if '/' in model_name else model_name - self._update_progress( - status_message=f"Loading {'vision' if self.is_vlm else 'text'} model... {model_display}" - ) - - print(f"\nLoading {'vision' if self.is_vlm else 'text'} model: {model_name}") - - # Set HF token if provided - if hf_token: - os.environ["HF_TOKEN"] = hf_token - - - # Branch based on model type - if self.is_vlm: - # Load vision model - returns (model, tokenizer) - self.model, self.tokenizer = FastVisionModel.from_pretrained( - model_name=model_name, - max_seq_length=max_seq_length, - dtype=None, # Auto-detect - load_in_4bit=load_in_4bit, - token=hf_token, - ) - logger.info("Loaded vision model") - else: - # Load text model - returns (model, tokenizer) - self.model, self.tokenizer = FastLanguageModel.from_pretrained( - model_name=model_name, - max_seq_length=max_seq_length, - dtype=None, # Auto-detect - load_in_4bit=load_in_4bit, - token=hf_token, - ) - logger.info("Loaded text model") - - if self.should_stop: - return False - - self._update_progress(status_message="Model loaded successfully") - print("Model loaded successfully") - return True - - except Exception as e: - logger.error(f"Error loading model: {e}") - self._update_progress(error=str(e), is_training=False) - return False - - def prepare_model_for_training(self, - use_lora: bool = True, - # Vision-specific LoRA parameters (only used if is_vlm=True) - finetune_vision_layers: bool = True, - finetune_language_layers: bool = True, - finetune_attention_modules: bool = True, - finetune_mlp_modules: bool = True, - # Standard LoRA parameters - target_modules: list = None, - lora_r: int = 16, - lora_alpha: int = 16, - lora_dropout: float = 0.0, - use_gradient_checkpointing: str = "unsloth", - use_rslora: bool = False, - use_loftq: bool = False) -> bool: - """ - Prepare model for training (with optional LoRA). - """ - try: - if self.model is None: - raise ValueError("Model not loaded. Call load_model() first.") - - - # Full finetuning mode - skip PEFT entirely - if not use_lora: - self._update_progress(status_message="Full finetuning mode - no LoRA adapters") - print("Full finetuning mode - training all parameters\n") - return True - - # LoRA/QLoRA mode - apply PEFT - if target_modules is None or (isinstance(target_modules, list) and len(target_modules) == 0): - target_modules = ["q_proj", "k_proj", "v_proj", "o_proj", - "gate_proj", "up_proj", "down_proj"] - - # Validate and normalize gradient_checkpointing - # Must be one of: True, False, or "unsloth" - if isinstance(use_gradient_checkpointing, str): - use_gradient_checkpointing = use_gradient_checkpointing.strip().lower() - if use_gradient_checkpointing == "" or use_gradient_checkpointing == "unsloth": - use_gradient_checkpointing = "unsloth" - elif use_gradient_checkpointing in ("true", "1", "yes"): - use_gradient_checkpointing = True - elif use_gradient_checkpointing in ("false", "0", "no"): - use_gradient_checkpointing = False - else: - # Invalid value, default to "unsloth" - logger.warning(f"Invalid gradient_checkpointing value: {use_gradient_checkpointing}, defaulting to 'unsloth'") - use_gradient_checkpointing = "unsloth" - elif use_gradient_checkpointing not in (True, False, "unsloth"): - # Invalid type or value, default to "unsloth" - logger.warning(f"Invalid gradient_checkpointing type/value: {use_gradient_checkpointing}, defaulting to 'unsloth'") - use_gradient_checkpointing = "unsloth" - - # Verify model is loaded - if self.model is None: - error_msg = "Model is None - model was not loaded properly" - logger.error(error_msg) - self._update_progress(error=error_msg) - return False - - # Check if model has the expected attributes - if not hasattr(self.model, 'config'): - error_msg = "Model does not have config attribute - model may not be loaded correctly" - logger.error(error_msg) - self._update_progress(error=error_msg) - return False - - print(f"Configuring LoRA adapters (r={lora_r}, alpha={lora_alpha})...\n") - print(f"Gradient checkpointing: {use_gradient_checkpointing} (type: {type(use_gradient_checkpointing).__name__})\n") - - # Branch based on vision vs text - if self.is_vlm: - # Vision model LoRA - print(f"Vision model LoRA configuration:") - print(f" - Finetune vision layers: {finetune_vision_layers}") - print(f" - Finetune language layers: {finetune_language_layers}") - print(f" - Finetune attention modules: {finetune_attention_modules}") - print(f" - Finetune MLP modules: {finetune_mlp_modules}\n") - - self.model = FastVisionModel.get_peft_model( - self.model, - finetune_vision_layers=finetune_vision_layers, - finetune_language_layers=finetune_language_layers, - finetune_attention_modules=finetune_attention_modules, - finetune_mlp_modules=finetune_mlp_modules, - r=lora_r, - target_modules=target_modules, - lora_alpha=lora_alpha, - lora_dropout=lora_dropout, - bias="none", - use_gradient_checkpointing=use_gradient_checkpointing, - random_state=3407, - use_rslora=use_rslora, - loftq_config={"loftq_bits": 4, "loftq_iter": 1} if use_loftq else None, - ) - else: - # Text model LoRA - print(f"Text model LoRA configuration:") - print(f" - Target modules: {target_modules}\n") - - self.model = FastLanguageModel.get_peft_model( - self.model, - r=lora_r, - target_modules=target_modules, - lora_alpha=lora_alpha, - lora_dropout=lora_dropout, - bias="none", - use_gradient_checkpointing=use_gradient_checkpointing, - random_state=3407, - use_rslora=use_rslora, - loftq_config={"loftq_bits": 4, "loftq_iter": 1} if use_loftq else None, - ) - - # Check if stopped during LoRA preparation - if self.should_stop: - print("Stopped during LoRA configuration\n") - return False - - self._update_progress(status_message="LoRA adapters configured") - print("LoRA adapters configured successfully\n") - return True - - except Exception as e: - import traceback - import sys - error_details = f"{type(e).__name__}: {str(e)}" if str(e) else f"{type(e).__name__} (no message)" - full_traceback = traceback.format_exc() - logger.error(f"Error preparing model: {error_details}") - logger.error(f"Full traceback:\n{full_traceback}") - print(f"\n[ERROR] Error preparing model: {error_details}", file=sys.stderr, flush=True) - print(f"[ERROR] Full traceback:\n{full_traceback}", file=sys.stderr, flush=True) - self._update_progress(error=error_details) - return False - - def load_and_format_dataset(self, - dataset_source: str, - format_type: str = "auto", - local_datasets: list = None) -> Optional[Dataset]: - """ - Load and prepare dataset for training - """ - try: - dataset = None - - if local_datasets: - # Load local datasets - all_data = [] - for dataset_file in local_datasets: - # dataset_file may already be an absolute path from routes/training.py - if os.path.isabs(dataset_file): - file_path = dataset_file - else: - # Fallback: try relative to assets/datasets - script_dir = Path(__file__).parent.parent - assets_datasets_dir = script_dir / "assets" / "datasets" - file_path = assets_datasets_dir / dataset_file - - if str(file_path).endswith('.json'): - with open(file_path, 'r', encoding='utf-8') as f: - data = json.load(f) - if isinstance(data, list): - all_data.extend(data) - else: - all_data.append(data) - elif str(file_path).endswith('.csv'): - df = pd.read_csv(file_path) - all_data.extend(df.to_dict('records')) - - if all_data: - dataset = Dataset.from_list(all_data) - - # Check if stopped during dataset loading - if self.should_stop: - print("Stopped during dataset loading\n") - return None - - self._update_progress(status_message=f"Loaded {len(all_data)} samples from local files") - print(f"Loaded {len(all_data)} samples from local files\n") - - elif dataset_source: - # Load from Hugging Face - dataset = load_dataset(dataset_source, split="train") - - # Check if stopped during dataset loading - if self.should_stop: - print("Stopped during dataset loading\n") - return None - - self._update_progress(status_message=f"Loaded dataset from HuggingFace: {dataset_source}") - print(f"Loaded dataset from Hugging Face: {dataset_source}\n") - - if dataset is None: - raise ValueError("No dataset provided") - - # Check if stopped before applying template - if self.should_stop: - print("Stopped before applying chat template\n") - return None - - # NEW: Use unified format_and_template_dataset - print(f"Formatting dataset with format_type='{format_type}'...\n") - - #breakpoint() - dataset_info = format_and_template_dataset( - dataset, - model_name=self.model_name, - tokenizer=self.tokenizer, # Works for both text and vision models - is_vlm=self.is_vlm, - format_type=format_type, # "auto", "alpaca", "chatml", "sharegpt" - dataset_name=dataset_source, - ) - - # Check if stopped during formatting - if self.should_stop: - print("Stopped during dataset formatting\n") - return None - - self._update_progress(status_message=f"Dataset formatted and ready for training") - print(f"Dataset formatted successfully\n") - return dataset_info - - except Exception as e: - logger.error(f"Error loading dataset: {e}") - self._update_progress(error=str(e)) - return None - - def start_training(self, - dataset: Dataset, - output_dir: str = "./outputs", - num_epochs: int = 3, - learning_rate: float = 5e-5, - batch_size: int = 2, - gradient_accumulation_steps: int = 4, - warmup_steps: int = None, - warmup_ratio: float = None, - max_steps: int = 0, - save_steps: int = 0, - weight_decay: float = 0.01, - random_seed: int = 3407, - packing: bool = False, - train_on_completions: bool = False, - enable_wandb: bool = False, - wandb_project: str = "unsloth-training", - wandb_token: str = None, - enable_tensorboard: bool = False, - tensorboard_dir: str = "runs", - **kwargs) -> bool: - """Start training in a separate thread""" - - if self.is_training: - logger.warning("Training already in progress") - return False - - - if self.model is None or self.tokenizer is None: - self._update_progress(error="Model not loaded") - return False - - # Start training in separate thread - self.training_thread = threading.Thread( - target=self._train_worker, - args=(dataset,), - kwargs={ - 'output_dir': output_dir, - 'num_epochs': num_epochs, - 'learning_rate': learning_rate, - 'batch_size': batch_size, - 'gradient_accumulation_steps': gradient_accumulation_steps, - 'warmup_steps': warmup_steps, - 'warmup_ratio': warmup_ratio, - 'max_steps': max_steps, - 'save_steps': save_steps, - 'weight_decay': weight_decay, - 'random_seed': random_seed, - 'packing': packing, - 'train_on_completions': train_on_completions, - 'enable_wandb': enable_wandb, - 'wandb_project': wandb_project, - 'wandb_token': wandb_token, - 'enable_tensorboard': enable_tensorboard, - 'tensorboard_dir': tensorboard_dir, - **kwargs - } - ) - - self.should_stop = False - self.training_thread.start() - return True - - def _train_worker(self, dataset: Dataset, **training_args): - """Worker function for training (runs in separate thread)""" - try: - self._update_progress(is_training=True, error=None) - - # Setup logging - if training_args.get('enable_wandb', False) and training_args.get('wandb_token'): - os.environ["WANDB_API_KEY"] = training_args['wandb_token'] - import wandb - wandb.init(project=training_args.get('wandb_project', 'unsloth-training')) - - # Create output directory - output_dir = training_args.get('output_dir', './outputs') - os.makedirs(output_dir, exist_ok=True) - - # ========== DATA COLLATOR SELECTION ========== - # Detect special model types - model_name_lower = self.model_name.lower() - is_deepseek_ocr = "deepseek" in model_name_lower and "ocr" in model_name_lower - - print("Configuring data collator...\n") - - data_collator = None # Default to built-in data collator - if is_deepseek_ocr: - # Special DeepSeek OCR collator - auto-install if needed - print("Detected DeepSeek OCR model\n") - # Ensure DeepSeek OCR module is installed - if not _ensure_deepseek_ocr_installed(): - error_msg = ( - "Failed to install DeepSeek OCR module. " - "Please install manually: " - "from huggingface_hub import snapshot_download; " - "snapshot_download('unsloth/DeepSeek-OCR', local_dir='deepseek_ocr')" - ) - logger.error(error_msg) - self._update_progress(error=error_msg, is_training=False) - return - - try: - from backend.data_utils import DeepSeekOCRDataCollator - - print("Configuring DeepSeek OCR data collator...\n") - FastVisionModel.for_training(self.model) - data_collator = DeepSeekOCRDataCollator( - tokenizer=self.tokenizer, - model=self.model, - image_size=640, - base_size=1024, - crop_mode=True, - train_on_responses_only=training_args.get('train_on_completions', False), - ) - print("DeepSeek OCR data collator configured successfully\n") - - except Exception as e: - logger.error(f"Failed to configure DeepSeek OCR collator: {e}") - error_msg = f"Error configuring DeepSeek OCR: {str(e)}" - self._update_progress(error=error_msg, is_training=False) - return - - elif self.is_vlm: - # Standard VLM collator - print("Using UnslothVisionDataCollator for vision model\n") - from unsloth.trainer import UnslothVisionDataCollator - - FastVisionModel.for_training(self.model) - data_collator = UnslothVisionDataCollator(self.model, self.tokenizer) - print("Vision data collator configured\n") - - # ========== TRAINING CONFIGURATION ========== - # Handle epochs vs max_steps properly - max_steps_val = training_args.get('max_steps', 0) - num_epochs_val = training_args.get('num_epochs', 3) - - # Handle warmup_steps vs warmup_ratio - warmup_steps_val = training_args.get('warmup_steps', None) - warmup_ratio_val = training_args.get('warmup_ratio', None) - - config_args = { - "per_device_train_batch_size": training_args.get('batch_size', 2), - "gradient_accumulation_steps": training_args.get('gradient_accumulation_steps', 4), - "num_train_epochs": training_args.get('num_epochs', 3), # Default to epochs - "learning_rate": training_args.get('learning_rate', 2e-4), - "fp16": not is_bfloat16_supported(), - "bf16": is_bfloat16_supported(), - "logging_steps": 1, - "weight_decay": training_args.get('weight_decay', 0.01), - "seed": training_args.get('random_seed', 3407), - "output_dir": output_dir, - "report_to": ["wandb"] if training_args.get('enable_wandb', False) else "none", - } - - # Add warmup parameter - use warmup_ratio if provided, otherwise warmup_steps - if warmup_ratio_val is not None: - config_args["warmup_ratio"] = warmup_ratio_val - print(f"Using warmup_ratio: {warmup_ratio_val}\n") - elif warmup_steps_val is not None: - config_args["warmup_steps"] = warmup_steps_val - print(f"Using warmup_steps: {warmup_steps_val}\n") - else: - # Default to warmup_steps if neither provided - config_args["warmup_steps"] = 5 - print(f"Using default warmup_steps: 5\n") - - # If max_steps is specified, use it instead of epochs - max_steps_val = training_args.get('max_steps', 0) - if max_steps_val and max_steps_val > 0: - del config_args["num_train_epochs"] # Remove epochs - config_args["max_steps"] = max_steps_val # Use steps instead - print(f"Training for {max_steps_val} steps\n") - else: - print(f"Training for {config_args['num_train_epochs']} epochs\n") - - # Add model-specific parameters - # Use optim and lr_scheduler_type from training_args if provided, otherwise use defaults - optim_value = training_args.get('optim', "adamw_8bit") - lr_scheduler_type_value = training_args.get('lr_scheduler_type', "linear") - - if self.is_vlm: - # Vision-specific config - print("Configuring vision model training parameters\n") - # Use provided values or defaults for vision models - optim_value = training_args.get('optim', "adamw_torch_fused") - lr_scheduler_type_value = training_args.get('lr_scheduler_type', "cosine") - config_args.update({ - "optim": optim_value, - "lr_scheduler_type": lr_scheduler_type_value, - "gradient_checkpointing": True, - "gradient_checkpointing_kwargs": {"use_reentrant": False}, - "max_grad_norm": 0.3, # Recommended for vision models - "remove_unused_columns": False, - "dataset_text_field": "", - "dataset_kwargs": {"skip_prepare_dataset": True}, - "max_length": training_args.get('max_seq_length', 2048), - }) - else: - print("Configuring text model training parameters\n") - config_args.update({ - "optim": optim_value, - "lr_scheduler_type": lr_scheduler_type_value, - "dataset_text_field": "text", - }) - - # Only add packing for text models (not DeepSeek OCR which is VLM) - if not is_deepseek_ocr: - packing_enabled = training_args.get('packing', False) - config_args["packing"] = packing_enabled - print(f"Sequence packing: {'enabled' if packing_enabled else 'disabled'}\n") - - print(f"The configuration is: {config_args}") - - print("Training configuration prepared\n") - # ========== TRAINER INITIALIZATION ========== - if self.is_vlm: - self.trainer = SFTTrainer( - model=self.model, - train_dataset=dataset['dataset'], - processing_class = self.tokenizer.tokenizer, - data_collator=data_collator, - args=SFTConfig(**config_args), - ) - else: - self.trainer = SFTTrainer( - model=self.model, - tokenizer=self.tokenizer, - train_dataset=dataset['dataset'], - data_collator=data_collator, - args=SFTConfig(**config_args), - ) - print("Trainer initialized\n") - - # ========== TRAIN ON RESPONSES ONLY ========== - # Determine if we should train on responses only - instruction_part = None - response_part = None - train_on_responses_enabled = training_args.get('train_on_completions', False) - - # DeepSeek OCR handles this internally in its collator, so skip - if train_on_responses_enabled and not (is_deepseek_ocr or dataset["final_format"].lower() == 'alpaca'): - try: - print("Configuring train on responses only...\n") - - # Get the template mapping for this model - model_name_lower = self.model_name.lower() - - if model_name_lower in MODEL_TO_TEMPLATE_MAPPER: - template_name = MODEL_TO_TEMPLATE_MAPPER[model_name_lower] - print(f"Detected template: {template_name}\n") - - if template_name in TEMPLATE_TO_RESPONSES_MAPPER: - instruction_part = TEMPLATE_TO_RESPONSES_MAPPER[template_name]["instruction"] - response_part = TEMPLATE_TO_RESPONSES_MAPPER[template_name]["response"] - - print(f"Instruction marker: {instruction_part[:50]}...\n") - print(f"Response marker: {response_part[:50]}...\n") - else: - print(f"No response mapping found for template: {template_name}\n") - train_on_responses_enabled = False - else: - print(f"No template mapping found for model: {self.model_name}\n") - train_on_responses_enabled = False - - except Exception as e: - logger.warning(f"Could not configure train on responses: {e}") - train_on_responses_enabled = False - - # Apply train on responses only if we have valid parts - if train_on_responses_enabled and instruction_part and response_part and not (is_deepseek_ocr or dataset["final_format"].lower() == 'alpaca'): - try: - from unsloth.chat_templates import train_on_responses_only - - self.trainer = train_on_responses_only( - self.trainer, - instruction_part=instruction_part, - response_part=response_part, - ) - print("Train on responses only configured successfully\n") - except Exception as e: - logger.warning(f"Failed to apply train on responses only: {e}") - train_on_responses_enabled = False - else: - if train_on_responses_enabled and is_deepseek_ocr: - print("Train on responses handled by DeepSeek OCR collator\n") - else: - print("Training on full sequences (including prompts)\n") - - # Add custom callback for progress tracking - from transformers import TrainerCallback - - class ProgressCallback(TrainerCallback): - def __init__(self, trainer_instance): - self.trainer_instance = trainer_instance - - def on_train_begin(self, args, state, control, **kwargs): - """Called at the beginning of training""" - pass - - def on_log(self, args, state, control, logs=None, **kwargs): - """Called when logging occurs""" - if logs: - # Get loss from either 'loss' or 'train_loss' key - loss_value = logs.get('loss', logs.get('train_loss', 0.0)) - self.trainer_instance._update_progress( - step=state.global_step, - epoch=round(state.epoch, 2) if state.epoch else 0, # Round epoch to 2 decimals - loss=loss_value, - learning_rate=logs.get('learning_rate', 0.0), - status_message="" # Clear status message so metrics show - ) - - def on_epoch_end(self, args, state, control, **kwargs): - """Called at the end of each epoch""" - self.trainer_instance._update_progress( - epoch=state.epoch, - step=state.global_step - ) - - def on_step_end(self, args, state, control, **kwargs): - """Called at the end of each step""" - # Check if we should stop training - if self.trainer_instance.should_stop: - print(f"Stop detected at step {state.global_step}\n") - control.should_training_stop = True - return control - - # ========== PROGRESS TRACKING ========== - progress_callback = ProgressCallback(self) - self.trainer.add_callback(progress_callback) - - num_samples = len(dataset["dataset"]) - batch_size = training_args.get('batch_size', 2) - grad_accum = training_args.get('gradient_accumulation_steps', 4) - num_epochs = training_args.get('num_epochs', 3) - max_steps_val = training_args.get('max_steps', 0) - - # Step 1: Calculate dataloader length (number of batches) - len_dataloader = math.ceil(num_samples / batch_size) - - # Step 2: Calculate steps per epoch (following transformers logic) - num_update_steps_per_epoch = max( - len_dataloader // grad_accum + int(len_dataloader % grad_accum > 0), - 1 - ) - - # Step 3: Determine total steps based on max_steps or epochs - if max_steps_val and max_steps_val > 0: - # Use max_steps if specified - total_steps = max_steps_val - print(f"Progress tracking: {total_steps} steps (max_steps)\n") - else: - # Calculate from epochs - total_steps = num_update_steps_per_epoch * num_epochs - print(f"Progress tracking: {total_steps} steps ({num_epochs} epochs × {num_update_steps_per_epoch} steps/epoch)\n") - - self._update_progress(total_steps=total_steps) - - # ========== START TRAINING ========== - self._update_progress(status_message="Starting training...") - print("Starting training...\n") - self.trainer.train() - - # ========== SAVE MODEL ========== - self.trainer.save_model() - self.tokenizer.save_pretrained(output_dir) - print(f"\nTraining completed! Model saved to {output_dir}\n") - - self._update_progress( - is_training=False, - is_completed=True, - #status_message=status_msg - status_message=f"Training completed! Model saved to {output_dir}", - ) - - except Exception as e: - logger.error(f"Training error: {e}") - self._update_progress(is_training=False, error=str(e)) - - finally: - self.is_training = False - - def stop_training(self): - """Stop ongoing training""" - print("\nStopping training...") - self.should_stop = True - self.is_training = False - # Clear the status message so timer doesn't show stale status - self._update_progress(is_training=False, status_message="") - - # If trainer exists, try to stop it gracefully - if self.trainer: - try: - # The callback will catch should_stop flag and stop the training loop - print("Training will stop at next step...\n") - except Exception as e: - logger.error(f"Error stopping trainer: {e}") - - def get_training_progress(self) -> TrainingProgress: - """Get current training progress""" - with self._lock: - return self.training_progress - - def cleanup(self): - """Cleanup resources""" - if self.trainer: - self.trainer = None - if self.model: - self.model = None - if self.tokenizer: - self.tokenizer = None - - # Clear GPU memory - if torch.cuda.is_available(): - torch.cuda.empty_cache() - - -def _ensure_deepseek_ocr_installed(): - """ - Auto-install DeepSeek OCR module if not available. - Downloads from HuggingFace hub as a local module. - - Returns: - bool: True if available (either already installed or just installed) - """ - try: - # Try importing to see if already available - from deepseek_ocr.modeling_deepseekocr import format_messages - logger.info("DeepSeek OCR module already available") - return True - except ImportError: - pass - - try: - logger.info("DeepSeek OCR module not found. Auto-installing from HuggingFace...") - print("\n Downloading DeepSeek OCR module from HuggingFace...\n") - - from huggingface_hub import snapshot_download - import sys - import os - - # Get the script directory to install locally - script_dir = os.path.dirname(os.path.abspath(__file__)) - parent_dir = os.path.dirname(script_dir) # Go up to project root - - # Download to project root as 'deepseek_ocr' folder - local_dir = os.path.join(parent_dir, "deepseek_ocr") - - snapshot_download( - "unsloth/DeepSeek-OCR", - local_dir=local_dir, - local_dir_use_symlinks=False - ) - - # Add to sys.path if not already there - if parent_dir not in sys.path: - sys.path.insert(0, parent_dir) - - # Try importing again - from deepseek_ocr.modeling_deepseekocr import format_messages - - logger.info("DeepSeek OCR module installed successfully") - print("DeepSeek OCR module installed successfully!\n") - return True - - except Exception as e: - logger.error(f"Failed to install DeepSeek OCR module: {e}") - print(f"\n❌ Failed to install DeepSeek OCR module: {e}\n") - return False - -# Global trainer instance -_trainer_instance = None - -def get_trainer() -> UnslothTrainer: - """Get global trainer instance""" - global _trainer_instance - if _trainer_instance is None: - _trainer_instance = UnslothTrainer() - return _trainer_instance diff --git a/studio/backend/backend/training.py b/studio/backend/backend/training.py deleted file mode 100644 index 02e96128eb..0000000000 --- a/studio/backend/backend/training.py +++ /dev/null @@ -1,683 +0,0 @@ -""" -Training backend and UI integration -""" -import gradio as gr -import matplotlib.pyplot as plt -from typing import Dict, Any, Generator, Tuple -import logging - -from .trainer import get_trainer, TrainingProgress - -logger = logging.getLogger(__name__) - -# Plot styling constants -PLOT_WIDTH = 8 # Inches -PLOT_HEIGHT = 3.5 # Inches - - -class TrainingBackend: - """ - Training orchestration and UI integration. - Handles both text and vision models, LoRA and full finetuning. - """ - - def __init__(self): - self.trainer = get_trainer() - - # Training Metrics - self.loss_history = [] - self.lr_history = [] - self.step_history = [] - self.current_theme = "light" - - self.trainer.add_progress_callback(self._on_progress_update) - - logger.info("TrainingBackend initialized") - - def _on_progress_update(self, progress: TrainingProgress): - """Callback for progress updates""" - if progress.step > 0 and progress.loss > 0: - self.loss_history.append(progress.loss) - self.lr_history.append(progress.learning_rate) - self.step_history.append(progress.step) - - def start_training(self, - # Model parameters - model_name: str, - training_type: str, # NEW: "LoRA/QLoRA" or "Full Finetuning" - hf_token: str, - load_in_4bit: bool, - max_seq_length: int, - - # Dataset parameters - hf_dataset: str, - local_datasets: list, - format_type: str, # CHANGED: was data_template - - # Training parameters - num_epochs: int, - learning_rate: str, - batch_size: int, - gradient_accumulation_steps: int, - warmup_steps: int, # May be None even without default - warmup_ratio: float, # May be None even without default - max_steps: int, - save_steps: int, - weight_decay: float, - random_seed: int, - packing: bool, - optim: str, - lr_scheduler_type: str, - - # LoRA parameters - use_lora: bool, # Should be derived from training_type - lora_r: int, - lora_alpha: int, - lora_dropout: float, - target_modules: list, - gradient_checkpointing: str, - use_rslora: bool, - use_loftq: bool, - train_on_completions: bool, - - # NEW: Vision-specific LoRA parameters - finetune_vision_layers: bool, - finetune_language_layers: bool, - finetune_attention_modules: bool, - finetune_mlp_modules: bool, - - # Logging parameters - enable_wandb: bool, - wandb_token: str, - wandb_project: str, - enable_tensorboard: bool, - tensorboard_dir: str) -> Generator[Tuple, None, None]: - """ - Start training - yields UI updates as generator. - - Yields: - Tuple of (start_btn_update, stop_btn_update, progress_visible, config_visible) - """ - try: - # Reset stop flag and clear history - self.trainer.should_stop = False - self.loss_history = [] - self.lr_history = [] - self.step_history = [] - import time - output_dir = f"./outputs/{model_name.replace('/', '_')}_{int(time.time())}" - - # NEW: Derive use_lora from training_type - use_lora_actual = (training_type == "LoRA/QLoRA") - if use_lora_actual: print("using Lora") - else: print("using full finetuning") - logger.info(f"Starting training - Type: {training_type}, Model: {model_name}") - - # Yield initial status - buttons toggle immediately - yield ( - gr.update(interactive=False), # Start button disabled - gr.update(interactive=True), # Stop button enabled - gr.update(visible=True), # Training progress visible - #gr.update(visible=False) # Config selection hidden - ) - - # ========== LOAD MODEL ========== - logger.info("Loading model...") - success = self.trainer.load_model( - model_name=model_name, - max_seq_length=max_seq_length, - load_in_4bit=load_in_4bit if use_lora_actual else False, # Only 4bit for LoRA - hf_token=hf_token if hf_token.strip() else None - ) - - if not success or self.trainer.should_stop: - logger.error("Failed to load model or stopped by user") - return - - # Capture if this is a vision model - #self.current_training_session['is_vlm'] = self.trainer.is_vlm - - yield ( - gr.update(interactive=False), - gr.update(interactive=True), - gr.update(visible=True), - #gr.update(visible=False) - ) - - # ========== PREPARE MODEL FOR TRAINING ========== - if use_lora_actual: - logger.info("Preparing model with LoRA...") - success = self.trainer.prepare_model_for_training( - use_lora=True, - # Vision-specific parameters - finetune_vision_layers=finetune_vision_layers, - finetune_language_layers=finetune_language_layers, - finetune_attention_modules=finetune_attention_modules, - finetune_mlp_modules=finetune_mlp_modules, - # Standard LoRA parameters - target_modules=target_modules, - lora_r=lora_r, - lora_alpha=lora_alpha, - lora_dropout=lora_dropout, - use_gradient_checkpointing=gradient_checkpointing, - use_rslora=use_rslora, - use_loftq=use_loftq - ) - else: - logger.info("Preparing model for full finetuning...") - success = self.trainer.prepare_model_for_training( - use_lora=False # Full finetuning - ) - - if not success or self.trainer.should_stop: - logger.error("Failed to prepare model or stopped by user") - return - - yield ( - gr.update(interactive=False), - gr.update(interactive=True), - gr.update(visible=True), - #gr.update(visible=False) - ) - - # ========== LOAD DATASET ========== - logger.info("Loading dataset...") - #breakpoint() - dataset = self.trainer.load_and_format_dataset( - dataset_source=hf_dataset if hf_dataset.strip() else None, - format_type=format_type, - local_datasets=local_datasets if local_datasets else None - ) - - if dataset is None or self.trainer.should_stop: - logger.error("Failed to load dataset or stopped by user") - return - - yield ( - gr.update(interactive=False), - gr.update(interactive=True), - gr.update(visible=True), - #gr.update(visible=False) - ) - - # ========== START TRAINING ========== - # Convert learning rate string to float - try: - lr_value = float(learning_rate) - except ValueError: - logger.error(f"Invalid learning rate: {learning_rate}") - self.trainer._update_progress( - error=f"Invalid learning rate: {learning_rate}", - is_training=False - ) - return - - logger.info("Starting training worker thread...") - success = self.trainer.start_training( - dataset=dataset, - #output_dir=f"./outputs/{model_name.replace('/', '_')}_{int(__import__('time').time())}", - output_dir=output_dir, - num_epochs=num_epochs, - learning_rate=lr_value, - batch_size=batch_size, - gradient_accumulation_steps=gradient_accumulation_steps, - warmup_steps=warmup_steps, - warmup_ratio=warmup_ratio, - max_steps=max_steps if max_steps > 0 else 0, - save_steps=save_steps if save_steps > 0 else 0, - weight_decay=weight_decay, - random_seed=random_seed, - packing=packing, - train_on_completions=train_on_completions, - enable_wandb=enable_wandb, - wandb_project=wandb_project, - wandb_token=wandb_token if wandb_token.strip() else None, - enable_tensorboard=enable_tensorboard, - tensorboard_dir=tensorboard_dir, - max_seq_length=max_seq_length, # Pass through for config - optim=optim, - lr_scheduler_type=lr_scheduler_type, - ) - - if not success: - logger.error("Failed to start training") - yield ( - gr.update(interactive=True), - gr.update(interactive=False), - gr.update(visible=False), - #gr.update(visible=True) - ) - - except Exception as e: - logger.error(f"Error in start_training: {e}", exc_info=True) - self.trainer._update_progress( - error=str(e), - is_training=False - ) - yield ( - gr.update(interactive=True), - gr.update(interactive=False), - gr.update(visible=False), - #gr.update(visible=True) - ) - - def stop_training(self) -> Tuple: - """ - Stop ongoing training. - - Returns: - Tuple of (start_btn_update, stop_btn_update, progress_visible, config_visible) - """ - try: - logger.info("Stopping training...") - self.trainer.stop_training() - - return ( - gr.update(interactive=True), # Start button enabled - gr.update(interactive=False), # Stop button disabled - gr.update(visible=False), # Training progress hidden - #gr.update(visible=True) # Config selection visible - ) - except Exception as e: - logger.error(f"Error stopping training: {e}") - return ( - gr.update(interactive=True), - gr.update(interactive=False), - gr.update(visible=False), - #gr.update(visible=True) - ) - - def get_training_status(self, theme: str = "light") -> Tuple[plt.Figure, gr.update, gr.update, gr.update]: - """ - Get current training status and loss plot. - - Args: - theme: "light" or "dark" for plot styling - - Returns: - Tuple of (plot, start_btn, stop_btn, progress_visible) - """ - - try: - progress = self.trainer.get_training_progress() - - # If not training and not completed, return no updates - if not (progress.is_training or progress.is_completed or progress.error): - return (None, gr.update(), gr.update(), gr.update()) - - # Generate plot - plot = self._create_loss_plot(progress, theme) - - # If completed or error, enable start button - if progress.is_completed or progress.error: - return ( - plot, - gr.update(interactive=True), # Start button enabled - gr.update(interactive=False), # Stop button disabled - gr.update(visible=True), # Training progress visible - ) - - # Still training - no button updates - return (plot, gr.update(), gr.update(), gr.update()) - - except Exception as e: - logger.error(f"Error getting training status: {e}") - return (None, gr.update(), gr.update(), gr.update()) - - def refresh_plot_for_theme(self, theme: str) -> plt.Figure: - """ - Refresh plot with new theme. - - Args: - theme: "light" or "dark" - - Returns: - Updated matplotlib figure - """ - if theme and isinstance(theme, str) and theme in ['light', 'dark']: - self.current_theme = theme - - # Always generate plot if we have loss history - if self.loss_history: - progress = self.trainer.get_training_progress() - return self._create_loss_plot(progress, self.current_theme) - - return None - - def is_training_active(self) -> bool: - """ - Check if training is currently active (from load_model start to completion/error). - - Returns: - True if training is in progress, False otherwise - """ - try: - progress = self.trainer.get_training_progress() - # Training is active if is_training is True - # Also check if we're in loading/preparation phase (status_message indicates activity) - is_active = progress.is_training - # Also consider it active if we have a status message indicating loading/preparation - # but haven't completed or errored yet - if not is_active and not progress.is_completed and not progress.error: - status = progress.status_message or "" - if any(keyword in status.lower() for keyword in ["loading", "preparing", "training"]): - is_active = True - return is_active - except Exception as e: - logger.error(f"Error checking training state: {e}") - return False - - def _create_loss_plot(self, progress: TrainingProgress, theme: str = "light") -> plt.Figure: - """ - Create training loss plot with theme-aware styling. - - Args: - progress: Current training progress - theme: "light" or "dark" - - Returns: - Matplotlib figure - """ - plt.close('all') - - # Theme-specific styling - LIGHT_STYLE = { - "facecolor": "#ffffff", - "grid_color": "#d1d5db", - "line": "#16b88a", - "text": "#1f2937", - "empty_text": "#6b7280" - } - DARK_STYLE = { - "facecolor": "#292929", - "grid_color": "#404040", - "line": "#4ade80", - "text": "#e5e7eb", - "empty_text": "#9ca3af" - } - - style = LIGHT_STYLE if theme == "light" else DARK_STYLE - - fig, ax = plt.subplots(figsize=(PLOT_WIDTH, PLOT_HEIGHT)) - fig.patch.set_facecolor(style["facecolor"]) - ax.set_facecolor(style["facecolor"]) - - if self.loss_history: - steps = self.step_history - losses = self.loss_history - scatter_color = "#60a5fa" - # Scatter plot for raw loss points - ax.scatter( - steps, - losses, - s=16, - alpha=0.6, - color=scatter_color, - linewidths=0, - label="Training Loss (raw)", - ) - - # Moving average line overlay (trailing window) - MA_WINDOW = 20 # adjust smoothing aggressiveness - window = min(MA_WINDOW, len(losses)) - - if window >= 2: - cumsum = [0.0] - for v in losses: - cumsum.append(cumsum[-1] + float(v)) - - ma = [] - for i in range(len(losses)): - start = max(0, i - window + 1) - denom = i - start + 1 - ma.append((cumsum[i + 1] - cumsum[start]) / denom) - - ax.plot( - steps, - ma, - color=style["line"], - linewidth=2.5, - alpha=0.95, - label=f"Moving Avg ({ma[-1]:.4f})", - ) - - leg = ax.legend(frameon=False, fontsize=9) - for t in leg.get_texts(): - t.set_color(style["text"]) - - ax.set_xlabel('Steps', fontsize=10, color=style["text"]) - ax.set_ylabel('Loss', fontsize=10, color=style["text"]) - - # Build status message for title - if progress.error: - title = f"Error: {progress.error}" - elif progress.is_completed: - title = f"Training completed! Final loss: {progress.loss:.4f}" - elif progress.status_message: - title = progress.status_message - elif progress.step > 0: - title = f"Epoch: {progress.epoch} | Step: {progress.step}/{progress.total_steps} | Loss: {progress.loss:.4f}" - else: - title = "Training Loss" - - ax.set_title(title, fontsize=11, fontweight='bold', - pad=10, color=style["text"]) - - # Style grid and spines - ax.grid(True, alpha=0.4, linestyle='--', color=style["grid_color"]) - ax.tick_params(colors=style["text"], which='both') - ax.spines['top'].set_visible(False) - ax.spines['right'].set_visible(False) - ax.spines['bottom'].set_color(style["text"]) - ax.spines['left'].set_color(style["text"]) - else: - display_msg = progress.status_message if progress.status_message else 'Waiting for training data...' - ax.text(0.5, 0.5, display_msg, - ha='center', va='center', fontsize=16, - color=style["empty_text"], - transform=ax.transAxes) - ax.set_xticks([]) - ax.set_yticks([]) - for spine in ax.spines.values(): - spine.set_visible(False) - - fig.tight_layout() - return fig - - def _transfer_to_inference_backend(self) -> bool: - """ - Transfer the trained model to InferenceBackend. - Called automatically when training completes. - """ - print("=" * 60) - print("DEBUG: _transfer_to_inference_backend() CALLED") - print("=" * 60) - - try: - from .inference import get_inference_backend - - session = self.current_training_session - - # Check if already transferred - if session.get('transferred', False): - print("DEBUG: Already transferred, returning True") - logger.info("Model already transferred, skipping") - return True - - # Validate session data - if not session.get('base_model_name') or not session.get('output_dir'): - logger.warning("Training session incomplete, cannot transfer") - logger.warning(f"Session data: {session}") - return False - - inference_backend = get_inference_backend() - - base_model_name = session['base_model_name'] - output_dir = session['output_dir'] - is_lora = session['is_lora'] - is_vlm = session['is_vlm'] - - logger.info(f"=" * 60) - logger.info(f"TRANSFERRING MODEL TO INFERENCE BACKEND") - logger.info(f"=" * 60) - logger.info(f" Base model: {base_model_name}") - logger.info(f" Output dir: {output_dir}") - logger.info(f" Is LoRA: {is_lora}") - logger.info(f" Is VLM: {is_vlm}") - - # Transfer the model object directly from trainer memory. - # If is_lora is True, self.trainer.model is a PeftModel (Base + Adapter). - # If is_lora is False, it is the finetuned Base Model. - inference_backend.models[base_model_name] = { - "model": self.trainer.model, - "tokenizer": self.trainer.tokenizer, - "is_vision": is_vlm, - "is_lora": is_lora, - "model_path": base_model_name, - "base_model": None, - "loaded_adapters": {}, - # Unsloth/PEFT training keeps the active adapter named 'default' in memory - "active_adapter": "default" if is_lora else None, - } - - # For vision models, also transfer processor - if is_vlm: - if hasattr(self.trainer, 'tokenizer'): - inference_backend.models[base_model_name]["processor"] = self.trainer.tokenizer - logger.info(" Transferred processor for vision model") - - # Load chat template info - inference_backend._load_chat_template_info(base_model_name) - - # If it was LoRA, register the output path. - # This ensures the Eval UI dropdown (which lists files) knows that - # the model currently in memory corresponds to this specific output directory. - if is_lora: - inference_backend.models[base_model_name]["last_trained_adapter"] = output_dir - logger.info(f"Marked trained LoRA adapter: {output_dir}") - - # Set as active model - inference_backend.active_model_name = base_model_name - logger.info(f"Set active model: {base_model_name}") - - return True - - except Exception as e: - logger.error(f"Error transferring model to inference backend: {e}") - import traceback - traceback.print_exc() - return False - - -# ========== GLOBAL INSTANCE ========== -_training_backend = None - -def get_training_backend() -> TrainingBackend: - """Get global training backend instance""" - global _training_backend - if _training_backend is None: - _training_backend = TrainingBackend() - return _training_backend - - -# ========== UI HANDLER CREATION ========== -def create_training_handlers(train_components: Dict[str, Any]) -> Dict[str, Any]: - """ - Create training event handlers for Gradio UI components. - - Args: - train_components: Dictionary of Gradio components from train page - - Returns: - Dictionary of handler functions - """ - backend = get_training_backend() - - def start_training_handler(*args): - """Handler for start training button - yields status updates""" - try: - # Extract parameters in the order they're passed from the UI - (model_name, training_type, hf_token, load_4bit, max_seq_length, - hf_dataset, local_datasets, format_type, - num_epochs, learning_rate, batch_size, gradient_accumulation_steps, - warmup_steps, warmup_ratio, max_steps, save_steps, weight_decay, random_seed, packing, - optim, lr_scheduler_type, - use_lora, lora_r, lora_alpha, lora_dropout, target_modules, - gradient_checkpointing, use_rslora, use_loftq, train_on_completions, - finetune_vision_layers, finetune_language_layers, - finetune_attention_modules, finetune_mlp_modules, - enable_wandb, wandb_token, wandb_project, - enable_tensorboard, tensorboard_dir) = args - - # Start training with correctly named parameters - this is a generator - for update_tuple in backend.start_training( - model_name=model_name, - training_type=training_type, - hf_token=hf_token, - load_in_4bit=load_4bit, - max_seq_length=max_seq_length, - hf_dataset=hf_dataset, - local_datasets=local_datasets, - format_type=format_type, - num_epochs=num_epochs, - learning_rate=learning_rate, - batch_size=batch_size, - gradient_accumulation_steps=gradient_accumulation_steps, - warmup_steps=warmup_steps, - warmup_ratio=warmup_ratio, - max_steps=max_steps, - save_steps=save_steps, - weight_decay=weight_decay, - random_seed=random_seed, - packing=packing, - optim=optim, - lr_scheduler_type=lr_scheduler_type, - use_lora=use_lora, - lora_r=lora_r, - lora_alpha=lora_alpha, - lora_dropout=lora_dropout, - target_modules=target_modules, - gradient_checkpointing=gradient_checkpointing, - use_rslora=use_rslora, - use_loftq=use_loftq, - train_on_completions=train_on_completions, - finetune_vision_layers=finetune_vision_layers, - finetune_language_layers=finetune_language_layers, - finetune_attention_modules=finetune_attention_modules, - finetune_mlp_modules=finetune_mlp_modules, - enable_wandb=enable_wandb, - wandb_token=wandb_token, - wandb_project=wandb_project, - enable_tensorboard=enable_tensorboard, - tensorboard_dir=tensorboard_dir - ): - # Yield each status update to Gradio - yield update_tuple - - except Exception as e: - logger.error(f"Error in start_training_handler: {e}", exc_info=True) - yield ( - gr.update(interactive=True), # Start button - gr.update(interactive=False), # Stop button - gr.update(visible=False), # Training progress - #gr.update(visible=True) # Config selection - ) - - def stop_training_handler(): - """Handler for stop training button""" - return backend.stop_training() - - def update_training_status(): - """Periodic update of training status and plot""" - return backend.get_training_status(backend.current_theme) - - def refresh_plot_for_theme(theme): - """Refresh plot with new theme""" - return backend.refresh_plot_for_theme(theme) - - return { - 'start_training': start_training_handler, - 'stop_training': stop_training_handler, - 'update_status': update_training_status, - 'refresh_plot': refresh_plot_for_theme - }