diff --git a/backend/models/__init__.py b/backend/models/__init__.py deleted file mode 100644 index 07836f2168..0000000000 --- a/backend/models/__init__.py +++ /dev/null @@ -1,37 +0,0 @@ -""" -Pydantic models for API request/response schemas -""" -from .training import ( - TrainingStartRequest, - TrainingStartResponse, - TrainingStatusResponse, - TrainingMetricsResponse, - TrainingProgressResponse, -) -from .models import ( - ModelSearchRequest, - ModelSearchResponse, - ModelListResponse, - ModelConfigResponse, - LoRAScanResponse, - LoRAInfo, - ModelInfo, -) - -__all__ = [ - # Training schemas - "TrainingStartRequest", - "TrainingStartResponse", - "TrainingStatusResponse", - "TrainingMetricsResponse", - "TrainingProgressResponse", - # Model management schemas - "ModelSearchRequest", - "ModelSearchResponse", - "ModelListResponse", - "ModelConfigResponse", - "LoRAScanResponse", - "LoRAInfo", - "ModelInfo", -] - diff --git a/backend/models/models.py b/backend/models/models.py deleted file mode 100644 index 9561bf3288..0000000000 --- a/backend/models/models.py +++ /dev/null @@ -1,56 +0,0 @@ -""" -Pydantic schemas for Model Management API -""" -from pydantic import BaseModel, Field -from typing import Optional, List, Dict, Any - - -class ModelSearchRequest(BaseModel): - """Request schema for searching HuggingFace models""" - query: str = Field(..., description="Search query") - hf_token: Optional[str] = Field(None, description="HuggingFace token for authenticated searches") - - -class ModelInfo(BaseModel): - """Model information""" - id: str = Field(..., description="Model identifier") - name: Optional[str] = Field(None, description="Display name") - description: Optional[str] = Field(None, description="Model description") - size: Optional[str] = Field(None, description="Model size") - is_vision: bool = Field(False, description="Whether model is a vision model") - is_lora: bool = Field(False, description="Whether model is a LoRA adapter") - - -class ModelSearchResponse(BaseModel): - """Response schema for model search""" - models: List[ModelInfo] = Field(default_factory=list, description="List of matching models") - total: int = Field(0, description="Total number of results") - - -class ModelListResponse(BaseModel): - """Response schema for listing available models""" - models: List[ModelInfo] = Field(default_factory=list, description="List of available models") - default_models: List[str] = Field(default_factory=list, description="List of default model IDs") - - -class ModelConfigResponse(BaseModel): - """Response schema for model configuration""" - model_name: str = Field(..., description="Model identifier") - config: Dict[str, Any] = Field(..., description="Model configuration dictionary") - is_vision: bool = Field(False, description="Whether model is a vision model") - is_lora: bool = Field(False, description="Whether model is a LoRA adapter") - base_model: Optional[str] = Field(None, description="Base model if this is a LoRA adapter") - - -class LoRAInfo(BaseModel): - """LoRA adapter information""" - display_name: str = Field(..., description="Display name for the LoRA") - adapter_path: str = Field(..., description="Path to the LoRA adapter") - base_model: Optional[str] = Field(None, description="Base model identifier") - - -class LoRAScanResponse(BaseModel): - """Response schema for scanning trained LoRA adapters""" - loras: List[LoRAInfo] = Field(default_factory=list, description="List of found LoRA adapters") - outputs_dir: str = Field(..., description="Directory that was scanned") - diff --git a/backend/models/training.py b/backend/models/training.py deleted file mode 100644 index 9105eaa7d3..0000000000 --- a/backend/models/training.py +++ /dev/null @@ -1,96 +0,0 @@ -""" -Pydantic schemas for Training API -""" -from pydantic import BaseModel, Field -from typing import Optional, List - - -class TrainingStartRequest(BaseModel): - """Request schema for starting training""" - # Model parameters - model_name: str = Field(..., description="Model identifier (e.g., 'unsloth/llama-3-8b-bnb-4bit')") - training_type: str = Field(..., description="Training type: 'LoRA/QLoRA' or 'Full Finetuning'") - hf_token: Optional[str] = Field(None, description="HuggingFace token") - load_in_4bit: bool = Field(True, description="Load model in 4-bit quantization") - max_seq_length: int = Field(2048, description="Maximum sequence length") - - # Dataset parameters - hf_dataset: Optional[str] = Field(None, description="HuggingFace dataset identifier") - local_datasets: List[str] = Field(default_factory=list, description="List of local dataset paths") - format_type: str = Field(..., description="Dataset format type") - - # Training parameters - num_epochs: int = Field(1, description="Number of training epochs") - learning_rate: str = Field("2e-4", description="Learning rate") - batch_size: int = Field(1, description="Batch size") - gradient_accumulation_steps: int = Field(1, description="Gradient accumulation steps") - warmup_steps: Optional[int] = Field(None, description="Warmup steps") - warmup_ratio: Optional[float] = Field(None, description="Warmup ratio") - max_steps: Optional[int] = Field(None, description="Maximum training steps") - save_steps: int = Field(100, description="Steps between checkpoints") - weight_decay: float = Field(0.01, description="Weight decay") - random_seed: int = Field(42, description="Random seed") - packing: bool = Field(False, description="Enable sequence packing") - optim: str = Field("adamw_8bit", description="Optimizer") - lr_scheduler_type: str = Field("linear", description="Learning rate scheduler type") - - # LoRA parameters - use_lora: bool = Field(True, description="Use LoRA (derived from training_type)") - lora_r: int = Field(16, description="LoRA rank") - lora_alpha: int = Field(16, description="LoRA alpha") - lora_dropout: float = Field(0.0, description="LoRA dropout") - target_modules: List[str] = Field(default_factory=list, description="Target modules for LoRA") - gradient_checkpointing: str = Field("", description="Gradient checkpointing setting") - use_rslora: bool = Field(False, description="Use RSLoRA") - use_loftq: bool = Field(False, description="Use LoftQ") - train_on_completions: bool = Field(False, description="Train on completions only") - - # Vision-specific LoRA parameters - finetune_vision_layers: bool = Field(False, description="Finetune vision layers") - finetune_language_layers: bool = Field(False, description="Finetune language layers") - finetune_attention_modules: bool = Field(False, description="Finetune attention modules") - finetune_mlp_modules: bool = Field(False, description="Finetune MLP modules") - - # Logging parameters - enable_wandb: bool = Field(False, description="Enable Weights & Biases logging") - wandb_token: Optional[str] = Field(None, description="W&B token") - wandb_project: Optional[str] = Field(None, description="W&B project name") - enable_tensorboard: bool = Field(False, description="Enable TensorBoard logging") - tensorboard_dir: Optional[str] = Field(None, description="TensorBoard directory") - - -class TrainingStartResponse(BaseModel): - """Response schema for training start""" - status: str = Field(..., description="Status: 'started' or 'error'") - job_id: Optional[str] = Field(None, description="Training job ID") - message: str = Field(..., description="Status message") - error: Optional[str] = Field(None, description="Error message if status is 'error'") - - -class TrainingStatusResponse(BaseModel): - """Response schema for training status""" - status: str = Field(..., description="Status: 'idle', 'preparing', 'training', 'stopping', 'error'") - is_active: bool = Field(..., description="Whether training is currently active (actual training running)") - message: str = Field(..., description="Status message") - current_step: Optional[int] = Field(None, description="Current training step") - total_steps: Optional[int] = Field(None, description="Total training steps") - - -class TrainingMetricsResponse(BaseModel): - """Response schema for training metrics""" - loss_history: List[float] = Field(default_factory=list, description="Loss values") - lr_history: List[float] = Field(default_factory=list, description="Learning rate values") - step_history: List[int] = Field(default_factory=list, description="Step numbers") - current_loss: Optional[float] = Field(None, description="Current loss value") - current_lr: Optional[float] = Field(None, description="Current learning rate") - current_step: Optional[int] = Field(None, description="Current step") - - -class TrainingProgressResponse(BaseModel): - """Response schema for training progress updates""" - step: int = Field(..., description="Current step") - loss: float = Field(..., description="Current loss") - learning_rate: float = Field(..., description="Current learning rate") - status_message: str = Field(..., description="Status message") - progress_percent: Optional[float] = Field(None, description="Progress percentage") - diff --git a/frontend/.gitkeep b/frontend/.gitkeep deleted file mode 100644 index e69de29bb2..0000000000 diff --git a/backend/assets/datasets/alpaca_unsloth.json b/studio/backend/assets/datasets/alpaca_unsloth.json similarity index 100% rename from backend/assets/datasets/alpaca_unsloth.json rename to studio/backend/assets/datasets/alpaca_unsloth.json diff --git a/backend/auth/.gitkeep b/studio/backend/auth/.gitkeep similarity index 100% rename from backend/auth/.gitkeep rename to studio/backend/auth/.gitkeep diff --git a/backend/auth/__init__.py b/studio/backend/auth/__init__.py similarity index 100% rename from backend/auth/__init__.py rename to studio/backend/auth/__init__.py diff --git a/studio/backend/backend/__init__.py b/studio/backend/backend/__init__.py new file mode 100644 index 0000000000..33a28c425c --- /dev/null +++ b/studio/backend/backend/__init__.py @@ -0,0 +1,44 @@ +""" +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 new file mode 100644 index 0000000000..b1557624ab --- /dev/null +++ b/studio/backend/backend/export.py @@ -0,0 +1,506 @@ +# 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 new file mode 100644 index 0000000000..117487cb79 --- /dev/null +++ b/studio/backend/backend/inference.py @@ -0,0 +1,1212 @@ +""" +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 new file mode 100644 index 0000000000..85bdd2cf42 --- /dev/null +++ b/studio/backend/backend/model_config.py @@ -0,0 +1,704 @@ +""" +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/backend/utils/paths/path_utils.py b/studio/backend/backend/path_utils.py similarity index 100% rename from backend/utils/paths/path_utils.py rename to studio/backend/backend/path_utils.py diff --git a/studio/backend/backend/trainer.py b/studio/backend/backend/trainer.py new file mode 100644 index 0000000000..ce66e8ea6f --- /dev/null +++ b/studio/backend/backend/trainer.py @@ -0,0 +1,872 @@ +""" +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 new file mode 100644 index 0000000000..02e96128eb --- /dev/null +++ b/studio/backend/backend/training.py @@ -0,0 +1,683 @@ +""" +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 + } diff --git a/backend/core/__init__.py b/studio/backend/core/__init__.py similarity index 100% rename from backend/core/__init__.py rename to studio/backend/core/__init__.py diff --git a/backend/core/export/__init__.py b/studio/backend/core/export/__init__.py similarity index 100% rename from backend/core/export/__init__.py rename to studio/backend/core/export/__init__.py diff --git a/backend/core/export/export.py b/studio/backend/core/export/export.py similarity index 100% rename from backend/core/export/export.py rename to studio/backend/core/export/export.py diff --git a/backend/core/inference/__init__.py b/studio/backend/core/inference/__init__.py similarity index 100% rename from backend/core/inference/__init__.py rename to studio/backend/core/inference/__init__.py diff --git a/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py similarity index 100% rename from backend/core/inference/inference.py rename to studio/backend/core/inference/inference.py diff --git a/backend/core/training/__init__.py b/studio/backend/core/training/__init__.py similarity index 100% rename from backend/core/training/__init__.py rename to studio/backend/core/training/__init__.py diff --git a/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py similarity index 100% rename from backend/core/training/trainer.py rename to studio/backend/core/training/trainer.py diff --git a/backend/core/training/training.py b/studio/backend/core/training/training.py similarity index 100% rename from backend/core/training/training.py rename to studio/backend/core/training/training.py diff --git a/backend/loggers/.gitkeep b/studio/backend/loggers/.gitkeep similarity index 100% rename from backend/loggers/.gitkeep rename to studio/backend/loggers/.gitkeep diff --git a/backend/loggers/__init__.py b/studio/backend/loggers/__init__.py similarity index 100% rename from backend/loggers/__init__.py rename to studio/backend/loggers/__init__.py diff --git a/backend/main.py b/studio/backend/main.py similarity index 100% rename from backend/main.py rename to studio/backend/main.py diff --git a/backend/requirements.txt b/studio/backend/requirements.txt similarity index 100% rename from backend/requirements.txt rename to studio/backend/requirements.txt diff --git a/backend/models/.gitkeep b/studio/backend/routes/.gitkeep similarity index 100% rename from backend/models/.gitkeep rename to studio/backend/routes/.gitkeep diff --git a/backend/routes/__init__.py b/studio/backend/routes/__init__.py similarity index 100% rename from backend/routes/__init__.py rename to studio/backend/routes/__init__.py diff --git a/backend/routes/inference.py b/studio/backend/routes/inference.py similarity index 100% rename from backend/routes/inference.py rename to studio/backend/routes/inference.py diff --git a/backend/routes/models.py b/studio/backend/routes/models.py similarity index 100% rename from backend/routes/models.py rename to studio/backend/routes/models.py diff --git a/backend/routes/training.py b/studio/backend/routes/training.py similarity index 100% rename from backend/routes/training.py rename to studio/backend/routes/training.py diff --git a/backend/run.py b/studio/backend/run.py similarity index 100% rename from backend/run.py rename to studio/backend/run.py diff --git a/backend/routes/.gitkeep b/studio/backend/state/.gitkeep similarity index 100% rename from backend/routes/.gitkeep rename to studio/backend/state/.gitkeep diff --git a/backend/state/__init__.py b/studio/backend/state/__init__.py similarity index 100% rename from backend/state/__init__.py rename to studio/backend/state/__init__.py diff --git a/backend/state/.gitkeep b/studio/backend/utils/.gitkeep similarity index 100% rename from backend/state/.gitkeep rename to studio/backend/utils/.gitkeep diff --git a/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py similarity index 100% rename from backend/utils/datasets/dataset_utils.py rename to studio/backend/utils/datasets/dataset_utils.py diff --git a/backend/utils/paths/__init__.py b/studio/backend/utils/paths/__init__.py similarity index 100% rename from backend/utils/paths/__init__.py rename to studio/backend/utils/paths/__init__.py diff --git a/studio/backend/utils/paths/path_utils.py b/studio/backend/utils/paths/path_utils.py new file mode 100644 index 0000000000..7743952b6b --- /dev/null +++ b/studio/backend/utils/paths/path_utils.py @@ -0,0 +1,78 @@ +""" +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/backend/utils/utils.py b/studio/backend/utils/utils.py similarity index 100% rename from backend/utils/utils.py rename to studio/backend/utils/utils.py diff --git a/frontend/.gitignore b/studio/frontend/.gitignore similarity index 100% rename from frontend/.gitignore rename to studio/frontend/.gitignore diff --git a/backend/utils/.gitkeep b/studio/frontend/.gitkeep similarity index 100% rename from backend/utils/.gitkeep rename to studio/frontend/.gitkeep diff --git a/frontend/README.md b/studio/frontend/README.md similarity index 100% rename from frontend/README.md rename to studio/frontend/README.md diff --git a/frontend/biome.json b/studio/frontend/biome.json similarity index 100% rename from frontend/biome.json rename to studio/frontend/biome.json diff --git a/frontend/bun.lock b/studio/frontend/bun.lock similarity index 100% rename from frontend/bun.lock rename to studio/frontend/bun.lock diff --git a/frontend/components.json b/studio/frontend/components.json similarity index 100% rename from frontend/components.json rename to studio/frontend/components.json diff --git a/frontend/eslint.config.js b/studio/frontend/eslint.config.js similarity index 100% rename from frontend/eslint.config.js rename to studio/frontend/eslint.config.js diff --git a/frontend/index.html b/studio/frontend/index.html similarity index 100% rename from frontend/index.html rename to studio/frontend/index.html diff --git a/frontend/package.json b/studio/frontend/package.json similarity index 100% rename from frontend/package.json rename to studio/frontend/package.json diff --git a/frontend/public/Hellix font official/OTF/Hellix-SemiBold.otf b/studio/frontend/public/Hellix font official/OTF/Hellix-SemiBold.otf similarity index 100% rename from frontend/public/Hellix font official/OTF/Hellix-SemiBold.otf rename to studio/frontend/public/Hellix font official/OTF/Hellix-SemiBold.otf diff --git a/frontend/public/Hellix font official/TTF/Hellix-SemiBold.ttf b/studio/frontend/public/Hellix font official/TTF/Hellix-SemiBold.ttf similarity index 100% rename from frontend/public/Hellix font official/TTF/Hellix-SemiBold.ttf rename to studio/frontend/public/Hellix font official/TTF/Hellix-SemiBold.ttf diff --git a/frontend/public/Hellix font official/WEB/Hellix-SemiBold.woff b/studio/frontend/public/Hellix font official/WEB/Hellix-SemiBold.woff similarity index 100% rename from frontend/public/Hellix font official/WEB/Hellix-SemiBold.woff rename to studio/frontend/public/Hellix font official/WEB/Hellix-SemiBold.woff diff --git a/frontend/public/Hellix font official/WEB/Hellix-SemiBold.woff2 b/studio/frontend/public/Hellix font official/WEB/Hellix-SemiBold.woff2 similarity index 100% rename from frontend/public/Hellix font official/WEB/Hellix-SemiBold.woff2 rename to studio/frontend/public/Hellix font official/WEB/Hellix-SemiBold.woff2 diff --git a/frontend/public/Sloth emojis/241024 Sloth Drink Jus.png b/studio/frontend/public/Sloth emojis/241024 Sloth Drink Jus.png similarity index 100% rename from frontend/public/Sloth emojis/241024 Sloth Drink Jus.png rename to studio/frontend/public/Sloth emojis/241024 Sloth Drink Jus.png diff --git a/frontend/public/Sloth emojis/251008 Sloth Pin.png b/studio/frontend/public/Sloth emojis/251008 Sloth Pin.png similarity index 100% rename from frontend/public/Sloth emojis/251008 Sloth Pin.png rename to studio/frontend/public/Sloth emojis/251008 Sloth Pin.png diff --git a/frontend/public/Sloth emojis/FO2C6766BA42 Sloth Gift.png b/studio/frontend/public/Sloth emojis/FO2C6766BA42 Sloth Gift.png similarity index 100% rename from frontend/public/Sloth emojis/FO2C6766BA42 Sloth Gift.png rename to studio/frontend/public/Sloth emojis/FO2C6766BA42 Sloth Gift.png diff --git a/frontend/public/Sloth emojis/FO71A40FA5581 Sloth and Llama.png b/studio/frontend/public/Sloth emojis/FO71A40FA5581 Sloth and Llama.png similarity index 100% rename from frontend/public/Sloth emojis/FO71A40FA5581 Sloth and Llama.png rename to studio/frontend/public/Sloth emojis/FO71A40FA5581 Sloth and Llama.png diff --git a/frontend/public/Sloth emojis/Large sloth Question mark.png b/studio/frontend/public/Sloth emojis/Large sloth Question mark.png similarity index 100% rename from frontend/public/Sloth emojis/Large sloth Question mark.png rename to studio/frontend/public/Sloth emojis/Large sloth Question mark.png diff --git a/frontend/public/Sloth emojis/Sloth loca pc.png b/studio/frontend/public/Sloth emojis/Sloth loca pc.png similarity index 100% rename from frontend/public/Sloth emojis/Sloth loca pc.png rename to studio/frontend/public/Sloth emojis/Sloth loca pc.png diff --git a/frontend/public/Sloth emojis/Sloth w Gameboy Confetti no Logo.png b/studio/frontend/public/Sloth emojis/Sloth w Gameboy Confetti no Logo.png similarity index 100% rename from frontend/public/Sloth emojis/Sloth w Gameboy Confetti no Logo.png rename to studio/frontend/public/Sloth emojis/Sloth w Gameboy Confetti no Logo.png diff --git a/frontend/public/Sloth emojis/Sloth w PC Confetti no Logo.png b/studio/frontend/public/Sloth emojis/Sloth w PC Confetti no Logo.png similarity index 100% rename from frontend/public/Sloth emojis/Sloth w PC Confetti no Logo.png rename to studio/frontend/public/Sloth emojis/Sloth w PC Confetti no Logo.png diff --git a/frontend/public/Sloth emojis/Sloth w PC no Logo.png b/studio/frontend/public/Sloth emojis/Sloth w PC no Logo.png similarity index 100% rename from frontend/public/Sloth emojis/Sloth w PC no Logo.png rename to studio/frontend/public/Sloth emojis/Sloth w PC no Logo.png diff --git a/frontend/public/Sloth emojis/UnSloth Eat GPU Mouth.png b/studio/frontend/public/Sloth emojis/UnSloth Eat GPU Mouth.png similarity index 100% rename from frontend/public/Sloth emojis/UnSloth Eat GPU Mouth.png rename to studio/frontend/public/Sloth emojis/UnSloth Eat GPU Mouth.png diff --git a/frontend/public/Sloth emojis/UnSloth Eat GPU.png b/studio/frontend/public/Sloth emojis/UnSloth Eat GPU.png similarity index 100% rename from frontend/public/Sloth emojis/UnSloth Eat GPU.png rename to studio/frontend/public/Sloth emojis/UnSloth Eat GPU.png diff --git a/frontend/public/Sloth emojis/UnSloth GPU Front square.png b/studio/frontend/public/Sloth emojis/UnSloth GPU Front square.png similarity index 100% rename from frontend/public/Sloth emojis/UnSloth GPU Front square.png rename to studio/frontend/public/Sloth emojis/UnSloth GPU Front square.png diff --git a/frontend/public/Sloth emojis/UnSloth Laptop.png b/studio/frontend/public/Sloth emojis/UnSloth Laptop.png similarity index 100% rename from frontend/public/Sloth emojis/UnSloth Laptop.png rename to studio/frontend/public/Sloth emojis/UnSloth Laptop.png diff --git a/frontend/public/Sloth emojis/UnSloth Sparkling large.png b/studio/frontend/public/Sloth emojis/UnSloth Sparkling large.png similarity index 100% rename from frontend/public/Sloth emojis/UnSloth Sparkling large.png rename to studio/frontend/public/Sloth emojis/UnSloth Sparkling large.png diff --git a/frontend/public/Sloth emojis/large sloth cheeky.png b/studio/frontend/public/Sloth emojis/large sloth cheeky.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth cheeky.png rename to studio/frontend/public/Sloth emojis/large sloth cheeky.png diff --git a/frontend/public/Sloth emojis/large sloth drink.png b/studio/frontend/public/Sloth emojis/large sloth drink.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth drink.png rename to studio/frontend/public/Sloth emojis/large sloth drink.png diff --git a/frontend/public/Sloth emojis/large sloth fire.png b/studio/frontend/public/Sloth emojis/large sloth fire.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth fire.png rename to studio/frontend/public/Sloth emojis/large sloth fire.png diff --git a/frontend/public/Sloth emojis/large sloth glasses.png b/studio/frontend/public/Sloth emojis/large sloth glasses.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth glasses.png rename to studio/frontend/public/Sloth emojis/large sloth glasses.png diff --git a/frontend/public/Sloth emojis/large sloth heart.png b/studio/frontend/public/Sloth emojis/large sloth heart.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth heart.png rename to studio/frontend/public/Sloth emojis/large sloth heart.png diff --git a/frontend/public/Sloth emojis/large sloth laugh.png b/studio/frontend/public/Sloth emojis/large sloth laugh.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth laugh.png rename to studio/frontend/public/Sloth emojis/large sloth laugh.png diff --git a/frontend/public/Sloth emojis/large sloth sad.png b/studio/frontend/public/Sloth emojis/large sloth sad.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth sad.png rename to studio/frontend/public/Sloth emojis/large sloth sad.png diff --git a/frontend/public/Sloth emojis/large sloth thumbs.png b/studio/frontend/public/Sloth emojis/large sloth thumbs.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth thumbs.png rename to studio/frontend/public/Sloth emojis/large sloth thumbs.png diff --git a/frontend/public/Sloth emojis/large sloth wave.png b/studio/frontend/public/Sloth emojis/large sloth wave.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth wave.png rename to studio/frontend/public/Sloth emojis/large sloth wave.png diff --git a/frontend/public/Sloth emojis/large sloth yay.png b/studio/frontend/public/Sloth emojis/large sloth yay.png similarity index 100% rename from frontend/public/Sloth emojis/large sloth yay.png rename to studio/frontend/public/Sloth emojis/large sloth yay.png diff --git a/frontend/public/Sloth emojis/sloth headphones.png b/studio/frontend/public/Sloth emojis/sloth headphones.png similarity index 100% rename from frontend/public/Sloth emojis/sloth headphones.png rename to studio/frontend/public/Sloth emojis/sloth headphones.png diff --git a/frontend/public/Sloth emojis/sloth huglove large.png b/studio/frontend/public/Sloth emojis/sloth huglove large.png similarity index 100% rename from frontend/public/Sloth emojis/sloth huglove large.png rename to studio/frontend/public/Sloth emojis/sloth huglove large.png diff --git a/frontend/public/Sloth emojis/sloth huglove large33.png b/studio/frontend/public/Sloth emojis/sloth huglove large33.png similarity index 100% rename from frontend/public/Sloth emojis/sloth huglove large33.png rename to studio/frontend/public/Sloth emojis/sloth huglove large33.png diff --git a/frontend/public/Sloth emojis/sloth magnify final(1).png b/studio/frontend/public/Sloth emojis/sloth magnify final(1).png similarity index 100% rename from frontend/public/Sloth emojis/sloth magnify final(1).png rename to studio/frontend/public/Sloth emojis/sloth magnify final(1).png diff --git a/frontend/public/Sloth emojis/sloth magnify final.png b/studio/frontend/public/Sloth emojis/sloth magnify final.png similarity index 100% rename from frontend/public/Sloth emojis/sloth magnify final.png rename to studio/frontend/public/Sloth emojis/sloth magnify final.png diff --git a/frontend/public/Sloth emojis/sloth on phone.png b/studio/frontend/public/Sloth emojis/sloth on phone.png similarity index 100% rename from frontend/public/Sloth emojis/sloth on phone.png rename to studio/frontend/public/Sloth emojis/sloth on phone.png diff --git a/frontend/public/Sloth emojis/sloth pc emoji.png b/studio/frontend/public/Sloth emojis/sloth pc emoji.png similarity index 100% rename from frontend/public/Sloth emojis/sloth pc emoji.png rename to studio/frontend/public/Sloth emojis/sloth pc emoji.png diff --git a/frontend/public/Sloth emojis/sloth pc square.png b/studio/frontend/public/Sloth emojis/sloth pc square.png similarity index 100% rename from frontend/public/Sloth emojis/sloth pc square.png rename to studio/frontend/public/Sloth emojis/sloth pc square.png diff --git a/frontend/public/Sloth emojis/sloth rounded.png b/studio/frontend/public/Sloth emojis/sloth rounded.png similarity index 100% rename from frontend/public/Sloth emojis/sloth rounded.png rename to studio/frontend/public/Sloth emojis/sloth rounded.png diff --git a/frontend/public/Sloth emojis/sloth shock large.png b/studio/frontend/public/Sloth emojis/sloth shock large.png similarity index 100% rename from frontend/public/Sloth emojis/sloth shock large.png rename to studio/frontend/public/Sloth emojis/sloth shock large.png diff --git a/frontend/public/Sloth emojis/sloth shy large.png b/studio/frontend/public/Sloth emojis/sloth shy large.png similarity index 100% rename from frontend/public/Sloth emojis/sloth shy large.png rename to studio/frontend/public/Sloth emojis/sloth shy large.png diff --git a/frontend/public/Sloth emojis/sloth sir large.png b/studio/frontend/public/Sloth emojis/sloth sir large.png similarity index 100% rename from frontend/public/Sloth emojis/sloth sir large.png rename to studio/frontend/public/Sloth emojis/sloth sir large.png diff --git a/frontend/public/Sloth emojis/sloth w pc transparent.png b/studio/frontend/public/Sloth emojis/sloth w pc transparent.png similarity index 100% rename from frontend/public/Sloth emojis/sloth w pc transparent.png rename to studio/frontend/public/Sloth emojis/sloth w pc transparent.png diff --git a/frontend/public/Sloth emojis/sloth with gameboy.png b/studio/frontend/public/Sloth emojis/sloth with gameboy.png similarity index 100% rename from frontend/public/Sloth emojis/sloth with gameboy.png rename to studio/frontend/public/Sloth emojis/sloth with gameboy.png diff --git a/frontend/public/fonts/Hellix-SemiBold.woff b/studio/frontend/public/fonts/Hellix-SemiBold.woff similarity index 100% rename from frontend/public/fonts/Hellix-SemiBold.woff rename to studio/frontend/public/fonts/Hellix-SemiBold.woff diff --git a/frontend/public/fonts/Hellix-SemiBold.woff2 b/studio/frontend/public/fonts/Hellix-SemiBold.woff2 similarity index 100% rename from frontend/public/fonts/Hellix-SemiBold.woff2 rename to studio/frontend/public/fonts/Hellix-SemiBold.woff2 diff --git a/frontend/public/huggingface.svg b/studio/frontend/public/huggingface.svg similarity index 100% rename from frontend/public/huggingface.svg rename to studio/frontend/public/huggingface.svg diff --git a/frontend/public/vite.svg b/studio/frontend/public/vite.svg similarity index 100% rename from frontend/public/vite.svg rename to studio/frontend/public/vite.svg diff --git a/frontend/src/app/app.tsx b/studio/frontend/src/app/app.tsx similarity index 100% rename from frontend/src/app/app.tsx rename to studio/frontend/src/app/app.tsx diff --git a/frontend/src/app/provider.tsx b/studio/frontend/src/app/provider.tsx similarity index 100% rename from frontend/src/app/provider.tsx rename to studio/frontend/src/app/provider.tsx diff --git a/frontend/src/app/router.tsx b/studio/frontend/src/app/router.tsx similarity index 100% rename from frontend/src/app/router.tsx rename to studio/frontend/src/app/router.tsx diff --git a/frontend/src/app/routes/__root.tsx b/studio/frontend/src/app/routes/__root.tsx similarity index 100% rename from frontend/src/app/routes/__root.tsx rename to studio/frontend/src/app/routes/__root.tsx diff --git a/frontend/src/app/routes/chat.tsx b/studio/frontend/src/app/routes/chat.tsx similarity index 100% rename from frontend/src/app/routes/chat.tsx rename to studio/frontend/src/app/routes/chat.tsx diff --git a/frontend/src/app/routes/grid-test.tsx b/studio/frontend/src/app/routes/grid-test.tsx similarity index 100% rename from frontend/src/app/routes/grid-test.tsx rename to studio/frontend/src/app/routes/grid-test.tsx diff --git a/frontend/src/app/routes/home.tsx b/studio/frontend/src/app/routes/home.tsx similarity index 100% rename from frontend/src/app/routes/home.tsx rename to studio/frontend/src/app/routes/home.tsx diff --git a/frontend/src/app/routes/onboarding.tsx b/studio/frontend/src/app/routes/onboarding.tsx similarity index 100% rename from frontend/src/app/routes/onboarding.tsx rename to studio/frontend/src/app/routes/onboarding.tsx diff --git a/frontend/src/app/routes/studio.tsx b/studio/frontend/src/app/routes/studio.tsx similarity index 100% rename from frontend/src/app/routes/studio.tsx rename to studio/frontend/src/app/routes/studio.tsx diff --git a/frontend/src/assets/react.svg b/studio/frontend/src/assets/react.svg similarity index 100% rename from frontend/src/assets/react.svg rename to studio/frontend/src/assets/react.svg diff --git a/frontend/src/components/assistant-ui/attachment.tsx b/studio/frontend/src/components/assistant-ui/attachment.tsx similarity index 100% rename from frontend/src/components/assistant-ui/attachment.tsx rename to studio/frontend/src/components/assistant-ui/attachment.tsx diff --git a/frontend/src/components/assistant-ui/markdown-text.tsx b/studio/frontend/src/components/assistant-ui/markdown-text.tsx similarity index 100% rename from frontend/src/components/assistant-ui/markdown-text.tsx rename to studio/frontend/src/components/assistant-ui/markdown-text.tsx diff --git a/frontend/src/components/assistant-ui/model-selector.tsx b/studio/frontend/src/components/assistant-ui/model-selector.tsx similarity index 100% rename from frontend/src/components/assistant-ui/model-selector.tsx rename to studio/frontend/src/components/assistant-ui/model-selector.tsx diff --git a/frontend/src/components/assistant-ui/reasoning.tsx b/studio/frontend/src/components/assistant-ui/reasoning.tsx similarity index 100% rename from frontend/src/components/assistant-ui/reasoning.tsx rename to studio/frontend/src/components/assistant-ui/reasoning.tsx diff --git a/frontend/src/components/assistant-ui/thread.tsx b/studio/frontend/src/components/assistant-ui/thread.tsx similarity index 100% rename from frontend/src/components/assistant-ui/thread.tsx rename to studio/frontend/src/components/assistant-ui/thread.tsx diff --git a/frontend/src/components/assistant-ui/tool-fallback.tsx b/studio/frontend/src/components/assistant-ui/tool-fallback.tsx similarity index 100% rename from frontend/src/components/assistant-ui/tool-fallback.tsx rename to studio/frontend/src/components/assistant-ui/tool-fallback.tsx diff --git a/frontend/src/components/assistant-ui/tooltip-icon-button.tsx b/studio/frontend/src/components/assistant-ui/tooltip-icon-button.tsx similarity index 100% rename from frontend/src/components/assistant-ui/tooltip-icon-button.tsx rename to studio/frontend/src/components/assistant-ui/tooltip-icon-button.tsx diff --git a/frontend/src/components/component-example.tsx b/studio/frontend/src/components/component-example.tsx similarity index 100% rename from frontend/src/components/component-example.tsx rename to studio/frontend/src/components/component-example.tsx diff --git a/frontend/src/components/example.tsx b/studio/frontend/src/components/example.tsx similarity index 100% rename from frontend/src/components/example.tsx rename to studio/frontend/src/components/example.tsx diff --git a/frontend/src/components/layout/dashboard-grid.tsx b/studio/frontend/src/components/layout/dashboard-grid.tsx similarity index 100% rename from frontend/src/components/layout/dashboard-grid.tsx rename to studio/frontend/src/components/layout/dashboard-grid.tsx diff --git a/frontend/src/components/layout/dashboard-layout.tsx b/studio/frontend/src/components/layout/dashboard-layout.tsx similarity index 100% rename from frontend/src/components/layout/dashboard-layout.tsx rename to studio/frontend/src/components/layout/dashboard-layout.tsx diff --git a/frontend/src/components/layout/index.ts b/studio/frontend/src/components/layout/index.ts similarity index 100% rename from frontend/src/components/layout/index.ts rename to studio/frontend/src/components/layout/index.ts diff --git a/frontend/src/components/navbar.tsx b/studio/frontend/src/components/navbar.tsx similarity index 100% rename from frontend/src/components/navbar.tsx rename to studio/frontend/src/components/navbar.tsx diff --git a/frontend/src/components/section-card.tsx b/studio/frontend/src/components/section-card.tsx similarity index 100% rename from frontend/src/components/section-card.tsx rename to studio/frontend/src/components/section-card.tsx diff --git a/frontend/src/components/ui/accordion.tsx b/studio/frontend/src/components/ui/accordion.tsx similarity index 100% rename from frontend/src/components/ui/accordion.tsx rename to studio/frontend/src/components/ui/accordion.tsx diff --git a/frontend/src/components/ui/alert-dialog.tsx b/studio/frontend/src/components/ui/alert-dialog.tsx similarity index 100% rename from frontend/src/components/ui/alert-dialog.tsx rename to studio/frontend/src/components/ui/alert-dialog.tsx diff --git a/frontend/src/components/ui/alert.tsx b/studio/frontend/src/components/ui/alert.tsx similarity index 100% rename from frontend/src/components/ui/alert.tsx rename to studio/frontend/src/components/ui/alert.tsx diff --git a/frontend/src/components/ui/animated-shiny-text.tsx b/studio/frontend/src/components/ui/animated-shiny-text.tsx similarity index 100% rename from frontend/src/components/ui/animated-shiny-text.tsx rename to studio/frontend/src/components/ui/animated-shiny-text.tsx diff --git a/frontend/src/components/ui/aspect-ratio.tsx b/studio/frontend/src/components/ui/aspect-ratio.tsx similarity index 100% rename from frontend/src/components/ui/aspect-ratio.tsx rename to studio/frontend/src/components/ui/aspect-ratio.tsx diff --git a/frontend/src/components/ui/avatar.tsx b/studio/frontend/src/components/ui/avatar.tsx similarity index 100% rename from frontend/src/components/ui/avatar.tsx rename to studio/frontend/src/components/ui/avatar.tsx diff --git a/frontend/src/components/ui/badge.tsx b/studio/frontend/src/components/ui/badge.tsx similarity index 100% rename from frontend/src/components/ui/badge.tsx rename to studio/frontend/src/components/ui/badge.tsx diff --git a/frontend/src/components/ui/breadcrumb.tsx b/studio/frontend/src/components/ui/breadcrumb.tsx similarity index 100% rename from frontend/src/components/ui/breadcrumb.tsx rename to studio/frontend/src/components/ui/breadcrumb.tsx diff --git a/frontend/src/components/ui/button.tsx b/studio/frontend/src/components/ui/button.tsx similarity index 100% rename from frontend/src/components/ui/button.tsx rename to studio/frontend/src/components/ui/button.tsx diff --git a/frontend/src/components/ui/calendar.tsx b/studio/frontend/src/components/ui/calendar.tsx similarity index 100% rename from frontend/src/components/ui/calendar.tsx rename to studio/frontend/src/components/ui/calendar.tsx diff --git a/frontend/src/components/ui/card.tsx b/studio/frontend/src/components/ui/card.tsx similarity index 100% rename from frontend/src/components/ui/card.tsx rename to studio/frontend/src/components/ui/card.tsx diff --git a/frontend/src/components/ui/chart.tsx b/studio/frontend/src/components/ui/chart.tsx similarity index 100% rename from frontend/src/components/ui/chart.tsx rename to studio/frontend/src/components/ui/chart.tsx diff --git a/frontend/src/components/ui/checkbox.tsx b/studio/frontend/src/components/ui/checkbox.tsx similarity index 100% rename from frontend/src/components/ui/checkbox.tsx rename to studio/frontend/src/components/ui/checkbox.tsx diff --git a/frontend/src/components/ui/collapsible.tsx b/studio/frontend/src/components/ui/collapsible.tsx similarity index 100% rename from frontend/src/components/ui/collapsible.tsx rename to studio/frontend/src/components/ui/collapsible.tsx diff --git a/frontend/src/components/ui/combobox.tsx b/studio/frontend/src/components/ui/combobox.tsx similarity index 100% rename from frontend/src/components/ui/combobox.tsx rename to studio/frontend/src/components/ui/combobox.tsx diff --git a/frontend/src/components/ui/command.tsx b/studio/frontend/src/components/ui/command.tsx similarity index 100% rename from frontend/src/components/ui/command.tsx rename to studio/frontend/src/components/ui/command.tsx diff --git a/frontend/src/components/ui/confetti.tsx b/studio/frontend/src/components/ui/confetti.tsx similarity index 100% rename from frontend/src/components/ui/confetti.tsx rename to studio/frontend/src/components/ui/confetti.tsx diff --git a/frontend/src/components/ui/context-menu.tsx b/studio/frontend/src/components/ui/context-menu.tsx similarity index 100% rename from frontend/src/components/ui/context-menu.tsx rename to studio/frontend/src/components/ui/context-menu.tsx diff --git a/frontend/src/components/ui/dialog.tsx b/studio/frontend/src/components/ui/dialog.tsx similarity index 100% rename from frontend/src/components/ui/dialog.tsx rename to studio/frontend/src/components/ui/dialog.tsx diff --git a/frontend/src/components/ui/dropdown-menu.tsx b/studio/frontend/src/components/ui/dropdown-menu.tsx similarity index 100% rename from frontend/src/components/ui/dropdown-menu.tsx rename to studio/frontend/src/components/ui/dropdown-menu.tsx diff --git a/frontend/src/components/ui/field.tsx b/studio/frontend/src/components/ui/field.tsx similarity index 100% rename from frontend/src/components/ui/field.tsx rename to studio/frontend/src/components/ui/field.tsx diff --git a/frontend/src/components/ui/hover-card.tsx b/studio/frontend/src/components/ui/hover-card.tsx similarity index 100% rename from frontend/src/components/ui/hover-card.tsx rename to studio/frontend/src/components/ui/hover-card.tsx diff --git a/frontend/src/components/ui/input-group.tsx b/studio/frontend/src/components/ui/input-group.tsx similarity index 100% rename from frontend/src/components/ui/input-group.tsx rename to studio/frontend/src/components/ui/input-group.tsx diff --git a/frontend/src/components/ui/input.tsx b/studio/frontend/src/components/ui/input.tsx similarity index 100% rename from frontend/src/components/ui/input.tsx rename to studio/frontend/src/components/ui/input.tsx diff --git a/frontend/src/components/ui/label.tsx b/studio/frontend/src/components/ui/label.tsx similarity index 100% rename from frontend/src/components/ui/label.tsx rename to studio/frontend/src/components/ui/label.tsx diff --git a/frontend/src/components/ui/light-rays.tsx b/studio/frontend/src/components/ui/light-rays.tsx similarity index 100% rename from frontend/src/components/ui/light-rays.tsx rename to studio/frontend/src/components/ui/light-rays.tsx diff --git a/frontend/src/components/ui/menubar.tsx b/studio/frontend/src/components/ui/menubar.tsx similarity index 100% rename from frontend/src/components/ui/menubar.tsx rename to studio/frontend/src/components/ui/menubar.tsx diff --git a/frontend/src/components/ui/navigation-menu.tsx b/studio/frontend/src/components/ui/navigation-menu.tsx similarity index 100% rename from frontend/src/components/ui/navigation-menu.tsx rename to studio/frontend/src/components/ui/navigation-menu.tsx diff --git a/frontend/src/components/ui/pagination.tsx b/studio/frontend/src/components/ui/pagination.tsx similarity index 100% rename from frontend/src/components/ui/pagination.tsx rename to studio/frontend/src/components/ui/pagination.tsx diff --git a/frontend/src/components/ui/popover.tsx b/studio/frontend/src/components/ui/popover.tsx similarity index 100% rename from frontend/src/components/ui/popover.tsx rename to studio/frontend/src/components/ui/popover.tsx diff --git a/frontend/src/components/ui/progress.tsx b/studio/frontend/src/components/ui/progress.tsx similarity index 100% rename from frontend/src/components/ui/progress.tsx rename to studio/frontend/src/components/ui/progress.tsx diff --git a/frontend/src/components/ui/radio-group.tsx b/studio/frontend/src/components/ui/radio-group.tsx similarity index 100% rename from frontend/src/components/ui/radio-group.tsx rename to studio/frontend/src/components/ui/radio-group.tsx diff --git a/frontend/src/components/ui/resizable.tsx b/studio/frontend/src/components/ui/resizable.tsx similarity index 100% rename from frontend/src/components/ui/resizable.tsx rename to studio/frontend/src/components/ui/resizable.tsx diff --git a/frontend/src/components/ui/scroll-area.tsx b/studio/frontend/src/components/ui/scroll-area.tsx similarity index 100% rename from frontend/src/components/ui/scroll-area.tsx rename to studio/frontend/src/components/ui/scroll-area.tsx diff --git a/frontend/src/components/ui/select.tsx b/studio/frontend/src/components/ui/select.tsx similarity index 100% rename from frontend/src/components/ui/select.tsx rename to studio/frontend/src/components/ui/select.tsx diff --git a/frontend/src/components/ui/separator.tsx b/studio/frontend/src/components/ui/separator.tsx similarity index 100% rename from frontend/src/components/ui/separator.tsx rename to studio/frontend/src/components/ui/separator.tsx diff --git a/frontend/src/components/ui/sheet.tsx b/studio/frontend/src/components/ui/sheet.tsx similarity index 100% rename from frontend/src/components/ui/sheet.tsx rename to studio/frontend/src/components/ui/sheet.tsx diff --git a/frontend/src/components/ui/sidebar.tsx b/studio/frontend/src/components/ui/sidebar.tsx similarity index 100% rename from frontend/src/components/ui/sidebar.tsx rename to studio/frontend/src/components/ui/sidebar.tsx diff --git a/frontend/src/components/ui/skeleton.tsx b/studio/frontend/src/components/ui/skeleton.tsx similarity index 100% rename from frontend/src/components/ui/skeleton.tsx rename to studio/frontend/src/components/ui/skeleton.tsx diff --git a/frontend/src/components/ui/slider.tsx b/studio/frontend/src/components/ui/slider.tsx similarity index 100% rename from frontend/src/components/ui/slider.tsx rename to studio/frontend/src/components/ui/slider.tsx diff --git a/frontend/src/components/ui/sonner.tsx b/studio/frontend/src/components/ui/sonner.tsx similarity index 100% rename from frontend/src/components/ui/sonner.tsx rename to studio/frontend/src/components/ui/sonner.tsx diff --git a/frontend/src/components/ui/sparkles-text.tsx b/studio/frontend/src/components/ui/sparkles-text.tsx similarity index 100% rename from frontend/src/components/ui/sparkles-text.tsx rename to studio/frontend/src/components/ui/sparkles-text.tsx diff --git a/frontend/src/components/ui/switch.tsx b/studio/frontend/src/components/ui/switch.tsx similarity index 100% rename from frontend/src/components/ui/switch.tsx rename to studio/frontend/src/components/ui/switch.tsx diff --git a/frontend/src/components/ui/table.tsx b/studio/frontend/src/components/ui/table.tsx similarity index 100% rename from frontend/src/components/ui/table.tsx rename to studio/frontend/src/components/ui/table.tsx diff --git a/frontend/src/components/ui/tabs.tsx b/studio/frontend/src/components/ui/tabs.tsx similarity index 100% rename from frontend/src/components/ui/tabs.tsx rename to studio/frontend/src/components/ui/tabs.tsx diff --git a/frontend/src/components/ui/textarea.tsx b/studio/frontend/src/components/ui/textarea.tsx similarity index 100% rename from frontend/src/components/ui/textarea.tsx rename to studio/frontend/src/components/ui/textarea.tsx diff --git a/frontend/src/components/ui/toggle-group.tsx b/studio/frontend/src/components/ui/toggle-group.tsx similarity index 100% rename from frontend/src/components/ui/toggle-group.tsx rename to studio/frontend/src/components/ui/toggle-group.tsx diff --git a/frontend/src/components/ui/toggle.tsx b/studio/frontend/src/components/ui/toggle.tsx similarity index 100% rename from frontend/src/components/ui/toggle.tsx rename to studio/frontend/src/components/ui/toggle.tsx diff --git a/frontend/src/components/ui/tooltip.tsx b/studio/frontend/src/components/ui/tooltip.tsx similarity index 100% rename from frontend/src/components/ui/tooltip.tsx rename to studio/frontend/src/components/ui/tooltip.tsx diff --git a/frontend/src/config/env.ts b/studio/frontend/src/config/env.ts similarity index 100% rename from frontend/src/config/env.ts rename to studio/frontend/src/config/env.ts diff --git a/frontend/src/config/training.ts b/studio/frontend/src/config/training.ts similarity index 100% rename from frontend/src/config/training.ts rename to studio/frontend/src/config/training.ts diff --git a/frontend/src/features/chat/adapter.ts b/studio/frontend/src/features/chat/adapter.ts similarity index 100% rename from frontend/src/features/chat/adapter.ts rename to studio/frontend/src/features/chat/adapter.ts diff --git a/frontend/src/features/chat/chat-page.tsx b/studio/frontend/src/features/chat/chat-page.tsx similarity index 100% rename from frontend/src/features/chat/chat-page.tsx rename to studio/frontend/src/features/chat/chat-page.tsx diff --git a/frontend/src/features/chat/chat-settings-sheet.tsx b/studio/frontend/src/features/chat/chat-settings-sheet.tsx similarity index 100% rename from frontend/src/features/chat/chat-settings-sheet.tsx rename to studio/frontend/src/features/chat/chat-settings-sheet.tsx diff --git a/frontend/src/features/chat/chat-top-bar.tsx b/studio/frontend/src/features/chat/chat-top-bar.tsx similarity index 100% rename from frontend/src/features/chat/chat-top-bar.tsx rename to studio/frontend/src/features/chat/chat-top-bar.tsx diff --git a/frontend/src/features/chat/db.ts b/studio/frontend/src/features/chat/db.ts similarity index 100% rename from frontend/src/features/chat/db.ts rename to studio/frontend/src/features/chat/db.ts diff --git a/frontend/src/features/chat/index.ts b/studio/frontend/src/features/chat/index.ts similarity index 100% rename from frontend/src/features/chat/index.ts rename to studio/frontend/src/features/chat/index.ts diff --git a/frontend/src/features/chat/runtime-provider.tsx b/studio/frontend/src/features/chat/runtime-provider.tsx similarity index 100% rename from frontend/src/features/chat/runtime-provider.tsx rename to studio/frontend/src/features/chat/runtime-provider.tsx diff --git a/frontend/src/features/chat/shared-composer.tsx b/studio/frontend/src/features/chat/shared-composer.tsx similarity index 100% rename from frontend/src/features/chat/shared-composer.tsx rename to studio/frontend/src/features/chat/shared-composer.tsx diff --git a/frontend/src/features/chat/thread-sidebar.tsx b/studio/frontend/src/features/chat/thread-sidebar.tsx similarity index 100% rename from frontend/src/features/chat/thread-sidebar.tsx rename to studio/frontend/src/features/chat/thread-sidebar.tsx diff --git a/frontend/src/features/chat/types.ts b/studio/frontend/src/features/chat/types.ts similarity index 100% rename from frontend/src/features/chat/types.ts rename to studio/frontend/src/features/chat/types.ts diff --git a/frontend/src/features/onboarding/components/splash-screen.tsx b/studio/frontend/src/features/onboarding/components/splash-screen.tsx similarity index 100% rename from frontend/src/features/onboarding/components/splash-screen.tsx rename to studio/frontend/src/features/onboarding/components/splash-screen.tsx diff --git a/frontend/src/features/onboarding/components/steps/dataset-step.tsx b/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx similarity index 100% rename from frontend/src/features/onboarding/components/steps/dataset-step.tsx rename to studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx diff --git a/frontend/src/features/onboarding/components/steps/hyperparameters-step.tsx b/studio/frontend/src/features/onboarding/components/steps/hyperparameters-step.tsx similarity index 100% rename from frontend/src/features/onboarding/components/steps/hyperparameters-step.tsx rename to studio/frontend/src/features/onboarding/components/steps/hyperparameters-step.tsx diff --git a/frontend/src/features/onboarding/components/steps/model-selection-step.tsx b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx similarity index 100% rename from frontend/src/features/onboarding/components/steps/model-selection-step.tsx rename to studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx diff --git a/frontend/src/features/onboarding/components/steps/model-type-step.tsx b/studio/frontend/src/features/onboarding/components/steps/model-type-step.tsx similarity index 100% rename from frontend/src/features/onboarding/components/steps/model-type-step.tsx rename to studio/frontend/src/features/onboarding/components/steps/model-type-step.tsx diff --git a/frontend/src/features/onboarding/components/steps/summary-step.tsx b/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx similarity index 100% rename from frontend/src/features/onboarding/components/steps/summary-step.tsx rename to studio/frontend/src/features/onboarding/components/steps/summary-step.tsx diff --git a/frontend/src/features/onboarding/components/wizard-content.tsx b/studio/frontend/src/features/onboarding/components/wizard-content.tsx similarity index 100% rename from frontend/src/features/onboarding/components/wizard-content.tsx rename to studio/frontend/src/features/onboarding/components/wizard-content.tsx diff --git a/frontend/src/features/onboarding/components/wizard-footer.tsx b/studio/frontend/src/features/onboarding/components/wizard-footer.tsx similarity index 100% rename from frontend/src/features/onboarding/components/wizard-footer.tsx rename to studio/frontend/src/features/onboarding/components/wizard-footer.tsx diff --git a/frontend/src/features/onboarding/components/wizard-layout.tsx b/studio/frontend/src/features/onboarding/components/wizard-layout.tsx similarity index 100% rename from frontend/src/features/onboarding/components/wizard-layout.tsx rename to studio/frontend/src/features/onboarding/components/wizard-layout.tsx diff --git a/frontend/src/features/onboarding/components/wizard-sidebar.tsx b/studio/frontend/src/features/onboarding/components/wizard-sidebar.tsx similarity index 100% rename from frontend/src/features/onboarding/components/wizard-sidebar.tsx rename to studio/frontend/src/features/onboarding/components/wizard-sidebar.tsx diff --git a/frontend/src/features/onboarding/components/wizard-step-item.tsx b/studio/frontend/src/features/onboarding/components/wizard-step-item.tsx similarity index 100% rename from frontend/src/features/onboarding/components/wizard-step-item.tsx rename to studio/frontend/src/features/onboarding/components/wizard-step-item.tsx diff --git a/frontend/src/features/onboarding/index.ts b/studio/frontend/src/features/onboarding/index.ts similarity index 100% rename from frontend/src/features/onboarding/index.ts rename to studio/frontend/src/features/onboarding/index.ts diff --git a/frontend/src/features/studio/index.ts b/studio/frontend/src/features/studio/index.ts similarity index 100% rename from frontend/src/features/studio/index.ts rename to studio/frontend/src/features/studio/index.ts diff --git a/frontend/src/features/studio/sections/charts-content.tsx b/studio/frontend/src/features/studio/sections/charts-content.tsx similarity index 100% rename from frontend/src/features/studio/sections/charts-content.tsx rename to studio/frontend/src/features/studio/sections/charts-content.tsx diff --git a/frontend/src/features/studio/sections/charts-section.tsx b/studio/frontend/src/features/studio/sections/charts-section.tsx similarity index 100% rename from frontend/src/features/studio/sections/charts-section.tsx rename to studio/frontend/src/features/studio/sections/charts-section.tsx diff --git a/frontend/src/features/studio/sections/config-summary-section.tsx b/studio/frontend/src/features/studio/sections/config-summary-section.tsx similarity index 100% rename from frontend/src/features/studio/sections/config-summary-section.tsx rename to studio/frontend/src/features/studio/sections/config-summary-section.tsx diff --git a/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx similarity index 100% rename from frontend/src/features/studio/sections/dataset-section.tsx rename to studio/frontend/src/features/studio/sections/dataset-section.tsx diff --git a/frontend/src/features/studio/sections/model-section.tsx b/studio/frontend/src/features/studio/sections/model-section.tsx similarity index 100% rename from frontend/src/features/studio/sections/model-section.tsx rename to studio/frontend/src/features/studio/sections/model-section.tsx diff --git a/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx similarity index 100% rename from frontend/src/features/studio/sections/params-section.tsx rename to studio/frontend/src/features/studio/sections/params-section.tsx diff --git a/frontend/src/features/studio/sections/progress-section.tsx b/studio/frontend/src/features/studio/sections/progress-section.tsx similarity index 100% rename from frontend/src/features/studio/sections/progress-section.tsx rename to studio/frontend/src/features/studio/sections/progress-section.tsx diff --git a/frontend/src/features/studio/sections/training-section.tsx b/studio/frontend/src/features/studio/sections/training-section.tsx similarity index 100% rename from frontend/src/features/studio/sections/training-section.tsx rename to studio/frontend/src/features/studio/sections/training-section.tsx diff --git a/frontend/src/features/studio/studio-page.tsx b/studio/frontend/src/features/studio/studio-page.tsx similarity index 100% rename from frontend/src/features/studio/studio-page.tsx rename to studio/frontend/src/features/studio/studio-page.tsx diff --git a/frontend/src/features/studio/training-view.tsx b/studio/frontend/src/features/studio/training-view.tsx similarity index 100% rename from frontend/src/features/studio/training-view.tsx rename to studio/frontend/src/features/studio/training-view.tsx diff --git a/frontend/src/hooks/index.ts b/studio/frontend/src/hooks/index.ts similarity index 100% rename from frontend/src/hooks/index.ts rename to studio/frontend/src/hooks/index.ts diff --git a/frontend/src/hooks/use-mobile.ts b/studio/frontend/src/hooks/use-mobile.ts similarity index 100% rename from frontend/src/hooks/use-mobile.ts rename to studio/frontend/src/hooks/use-mobile.ts diff --git a/frontend/src/index.css b/studio/frontend/src/index.css similarity index 100% rename from frontend/src/index.css rename to studio/frontend/src/index.css diff --git a/frontend/src/lib/utils.ts b/studio/frontend/src/lib/utils.ts similarity index 100% rename from frontend/src/lib/utils.ts rename to studio/frontend/src/lib/utils.ts diff --git a/frontend/src/main.tsx b/studio/frontend/src/main.tsx similarity index 100% rename from frontend/src/main.tsx rename to studio/frontend/src/main.tsx diff --git a/frontend/src/stores/index.ts b/studio/frontend/src/stores/index.ts similarity index 100% rename from frontend/src/stores/index.ts rename to studio/frontend/src/stores/index.ts diff --git a/frontend/src/stores/training.ts b/studio/frontend/src/stores/training.ts similarity index 100% rename from frontend/src/stores/training.ts rename to studio/frontend/src/stores/training.ts diff --git a/frontend/src/types/index.ts b/studio/frontend/src/types/index.ts similarity index 100% rename from frontend/src/types/index.ts rename to studio/frontend/src/types/index.ts diff --git a/frontend/src/types/training.ts b/studio/frontend/src/types/training.ts similarity index 100% rename from frontend/src/types/training.ts rename to studio/frontend/src/types/training.ts diff --git a/frontend/src/utils/index.ts b/studio/frontend/src/utils/index.ts similarity index 100% rename from frontend/src/utils/index.ts rename to studio/frontend/src/utils/index.ts diff --git a/frontend/tsconfig.app.json b/studio/frontend/tsconfig.app.json similarity index 100% rename from frontend/tsconfig.app.json rename to studio/frontend/tsconfig.app.json diff --git a/frontend/tsconfig.json b/studio/frontend/tsconfig.json similarity index 100% rename from frontend/tsconfig.json rename to studio/frontend/tsconfig.json diff --git a/frontend/tsconfig.node.json b/studio/frontend/tsconfig.node.json similarity index 100% rename from frontend/tsconfig.node.json rename to studio/frontend/tsconfig.node.json diff --git a/frontend/vite.config.ts b/studio/frontend/vite.config.ts similarity index 100% rename from frontend/vite.config.ts rename to studio/frontend/vite.config.ts