From f52bddc23ff30a27e1c17500a62dcdc1c9db987c Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Fri, 13 Feb 2026 09:25:49 +0000 Subject: [PATCH] refactor: remove gradio dependency from training backend --- studio/backend/core/__init__.py | 3 +- studio/backend/core/training/__init__.py | 3 +- studio/backend/core/training/trainer.py | 4 +- studio/backend/core/training/training.py | 214 +++-------------------- 4 files changed, 28 insertions(+), 196 deletions(-) diff --git a/studio/backend/core/__init__.py b/studio/backend/core/__init__.py index 7b562c6cb0..a4a700d45b 100644 --- a/studio/backend/core/__init__.py +++ b/studio/backend/core/__init__.py @@ -6,7 +6,7 @@ Unified core module for Unsloth backend from .inference import InferenceBackend, get_inference_backend # Training -from .training import UnslothTrainer, get_trainer, TrainingBackend, get_training_backend, create_training_handlers, TrainingProgress +from .training import UnslothTrainer, get_trainer, TrainingBackend, get_training_backend, TrainingProgress # Configuration (from utils) from utils.models import is_vision_model, ModelConfig, scan_trained_loras, load_model_defaults, get_base_model_from_lora @@ -27,7 +27,6 @@ __all__ = [ 'get_trainer', 'get_training_backend', 'TrainingBackend', - 'create_training_handlers', 'TrainingProgress', # Config diff --git a/studio/backend/core/training/__init__.py b/studio/backend/core/training/__init__.py index 65bf4c3501..8fc2a6c721 100644 --- a/studio/backend/core/training/__init__.py +++ b/studio/backend/core/training/__init__.py @@ -2,7 +2,7 @@ Training submodule - Training backends and trainer classes """ from .trainer import UnslothTrainer, get_trainer, TrainingProgress -from .training import TrainingBackend, get_training_backend, create_training_handlers +from .training import TrainingBackend, get_training_backend __all__ = [ 'UnslothTrainer', @@ -10,5 +10,4 @@ __all__ = [ 'TrainingProgress', 'TrainingBackend', 'get_training_backend', - 'create_training_handlers', ] diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 1466a28674..7017b1b558 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -1,6 +1,6 @@ """ Unsloth Training Backend -Integrates Unsloth training capabilities with the Gradio UI +Integrates Unsloth training capabilities with the FastAPI backend """ import torch from utils.hardware import clear_gpu_cache @@ -46,7 +46,7 @@ class TrainingProgress: class UnslothTrainer: """ - Unsloth Training Backend for Gradio UI Integration + Unsloth Training Backend """ def __init__(self): diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 9bd44400ac..f671280eb8 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -1,9 +1,8 @@ """ -Training backend and UI integration +Training backend for FastAPI integration """ -import gradio as gr import matplotlib.pyplot as plt -from typing import Dict, Any, Generator, Tuple +from typing import Any, Generator, Tuple import logging from .trainer import get_trainer, TrainingProgress @@ -17,7 +16,7 @@ PLOT_HEIGHT = 3.5 # Inches class TrainingBackend: """ - Training orchestration and UI integration. + Training orchestration backend. Handles both text and vision models, LoRA and full finetuning. """ @@ -91,12 +90,12 @@ class TrainingBackend: wandb_token: str, wandb_project: str, enable_tensorboard: bool, - tensorboard_dir: str) -> Generator[Tuple, None, None]: + tensorboard_dir: str) -> bool: """ - Start training - yields UI updates as generator. + Start training. - Yields: - Tuple of (start_btn_update, stop_btn_update, progress_visible, config_visible) + Returns: + True if training started successfully, False otherwise. """ try: # Reset stop flag and clear history @@ -107,20 +106,12 @@ class TrainingBackend: import time output_dir = f"./outputs/{model_name.replace('/', '_')}_{int(time.time())}" - # NEW: Derive use_lora from training_type + # 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( @@ -132,17 +123,7 @@ class TrainingBackend: 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) - ) + return False # ========== PREPARE MODEL FOR TRAINING ========== if use_lora_actual: @@ -171,14 +152,7 @@ class TrainingBackend: 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) - ) + return False # ========== LOAD DATASET ========== logger.info("Loading dataset...") @@ -191,14 +165,7 @@ class TrainingBackend: 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) - ) + return False # ========== START TRAINING ========== # Convert learning rate string to float @@ -241,12 +208,9 @@ class TrainingBackend: 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) - ) + return False + + return True except Exception as e: logger.error(f"Error in start_training: {e}", exc_info=True) @@ -254,40 +218,24 @@ class TrainingBackend: error=str(e), is_training=False ) - yield ( - gr.update(interactive=True), - gr.update(interactive=False), - gr.update(visible=False), - #gr.update(visible=True) - ) + return False - def stop_training(self) -> Tuple: + def stop_training(self) -> bool: """ Stop ongoing training. Returns: - Tuple of (start_btn_update, stop_btn_update, progress_visible, config_visible) + True if training was successfully stopped. """ 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 - ) + return True 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) - ) + return False - def get_training_status(self, theme: str = "light") -> Tuple[plt.Figure, gr.update, gr.update, gr.update]: + def get_training_status(self, theme: str = "light") -> Tuple: """ Get current training status and loss plot. @@ -295,7 +243,7 @@ class TrainingBackend: theme: "light" or "dark" for plot styling Returns: - Tuple of (plot, start_btn, stop_btn, progress_visible) + Tuple of (plot, progress) """ try: @@ -303,26 +251,15 @@ class TrainingBackend: # 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()) + return (None, progress) # 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()) + return (plot, progress) except Exception as e: logger.error(f"Error getting training status: {e}") - return (None, gr.update(), gr.update(), gr.update()) + return (None, None) def refresh_plot_for_theme(self, theme: str) -> plt.Figure: """ @@ -578,106 +515,3 @@ def get_training_backend() -> TrainingBackend: 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 - }