""" Training backend for FastAPI integration """ import matplotlib.pyplot as plt from typing import Any, Generator, Tuple import logging import math from .trainer import get_trainer, TrainingProgress from utils.hardware import clear_gpu_cache logger = logging.getLogger(__name__) # Plot styling constants PLOT_WIDTH = 8 # Inches PLOT_HEIGHT = 3.5 # Inches class TrainingBackend: """ Training orchestration backend. 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.grad_norm_history = [] self.grad_norm_step_history = [] self.eval_loss_history = [] self.eval_step_history = [] self.eval_enabled = False 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) if progress.step >= 0 and progress.grad_norm is not None: try: grad_norm = float(progress.grad_norm) except (TypeError, ValueError): grad_norm = None if grad_norm is not None and math.isfinite(grad_norm): self.grad_norm_history.append(grad_norm) self.grad_norm_step_history.append(progress.step) if progress.eval_loss is not None: self.eval_loss_history.append(progress.eval_loss) self.eval_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, # Optional parameters custom_format_mapping: dict = None, subset: str = None, train_split: str = "train", eval_split: str = None, eval_steps: float = 0.01, is_dataset_multimodal: bool = False) -> bool: """ Start training. Returns: True if training started successfully, False otherwise. """ try: # Wait for any previous training thread to finish old_thread = getattr(self.trainer, "training_thread", None) if old_thread and old_thread.is_alive(): logger.info("Waiting for previous training thread to finish...") old_thread.join(timeout=30) # Explicitly free old SFTTrainer and CUDA resources before loading new model. # Without this, forked multiprocessing workers (num_proc tokenization) inherit # stale CUDA state from the previous run, causing extreme slowdowns or crashes. if self.trainer.trainer is not None: logger.info("Cleaning up previous SFTTrainer...") self.trainer.trainer = None if self.trainer.model is not None: self.trainer.model = None if self.trainer.tokenizer is not None: self.trainer.tokenizer = None # Flush all pending async CUDA ops so forked tokenization processes # don't inherit stale async state that causes pool join to hang. import torch as _torch if _torch.cuda.is_available(): _torch.cuda.synchronize() import gc gc.collect() clear_gpu_cache() # Reset stop flag and clear history self.trainer.should_stop = False self.trainer.save_on_stop = True self.loss_history = [] self.lr_history = [] self.step_history = [] self.grad_norm_history = [] self.grad_norm_step_history = [] self.eval_loss_history = [] self.eval_step_history = [] self.eval_enabled = False import time output_dir = f"./outputs/{model_name.replace('/', '_')}_{int(time.time())}" # 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}") # ========== 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, is_dataset_multimodal=is_dataset_multimodal, ) if not success or self.trainer.should_stop: logger.error("Failed to load model or stopped by user") return 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 False # ========== LOAD DATASET ========== logger.info("Loading dataset...") #breakpoint() dataset_result = 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, custom_format_mapping=custom_format_mapping, subset=subset, train_split=train_split, eval_split=eval_split, ) # Unpack: load_and_format_dataset returns (dataset, eval_dataset) if isinstance(dataset_result, tuple): dataset, eval_dataset = dataset_result else: dataset = dataset_result eval_dataset = None # If user set eval_steps to 0, disable evaluation entirely if eval_steps is not None and float(eval_steps) <= 0: eval_dataset = None # Track whether eval is enabled for status reporting self.eval_enabled = eval_dataset is not None if dataset is None or self.trainer.should_stop: logger.error("Failed to load dataset or stopped by user") return 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, eval_dataset=eval_dataset, eval_steps=eval_steps, 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, optim=optim, lr_scheduler_type=lr_scheduler_type, ) if not success: logger.error("Failed to start training") return False return 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 ) return False def stop_training(self, save: bool = True) -> bool: """ Stop ongoing training. Args: save: If True, save the model at the current checkpoint. Returns: True if training was successfully stopped. """ try: logger.info(f"Stopping training (save={save})...") self.trainer.stop_training(save=save) return True except Exception as e: logger.error(f"Error stopping training: {e}") return False def get_training_status(self, theme: str = "light") -> Tuple: """ Get current training status and loss plot. Args: theme: "light" or "dark" for plot styling Returns: Tuple of (plot, progress) """ 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, progress) # Generate plot plot = self._create_loss_plot(progress, theme) return (plot, progress) except Exception as e: logger.error(f"Error getting training status: {e}") return (None, None) 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: training_thread = getattr(self.trainer, "training_thread", None) if training_thread and training_thread.is_alive(): return True # Stop requested and worker already exited => inactive. # This allows UI to show stopped state + "Back to configuration". if getattr(self.trainer, "should_stop", False): return False 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 "" status_lower = status.lower() if any( keyword in status_lower for keyword in ["cancelled", "canceled", "stopped", "completed", "ready to train"] ): return False if any( keyword in status_lower for keyword in [ "loading", "preparing", "training", "configuring", "tokenizing", "starting", ] ): 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