diff --git a/.gitignore b/.gitignore index 044775e846..904e15f974 100755 --- a/.gitignore +++ b/.gitignore @@ -10,6 +10,8 @@ __pycache__/ # Virtual environments .venv/ +.venv_overlay/ +.venv_t5/ venv/ env/ environment.yaml diff --git a/setup.ps1 b/setup.ps1 index 8a2ec56033..e46a7522d2 100644 --- a/setup.ps1 +++ b/setup.ps1 @@ -685,12 +685,43 @@ $CuTag = Get-PytorchCudaTag Write-Host " Installing PyTorch with CUDA support ($CuTag)..." -ForegroundColor Cyan pip install torch torchvision torchaudio --index-url "https://download.pytorch.org/whl/$CuTag" 2>&1 | Out-Null +# Install Triton for Windows (enables torch.compile — without it training can hang) +Write-Host " Installing Triton for Windows..." -ForegroundColor Cyan +pip install "triton-windows<3.7" 2>&1 | Out-Null +Write-Host "[OK] Triton for Windows installed (enables torch.compile)" -ForegroundColor Green + # Ordered heavy dependency installation — shared cross-platform script Write-Host " Running ordered dependency installation..." -ForegroundColor Cyan python "$PSScriptRoot\install_python_stack.py" # Restore ErrorActionPreference after pip/python work $ErrorActionPreference = $prevEAP +# ── Pre-install transformers 5.x into .venv_t5/ ── +# Models like GLM-4.7-Flash need transformers>=5.2.0. Instead of pip-installing +# at runtime (slow, ~10-15s), we pre-install into a separate directory. +# The training subprocess just prepends .venv_t5/ to sys.path — instant switch. +Write-Host "" +Write-Host " Pre-installing transformers 5.x for newer model support..." -ForegroundColor Cyan +$VenvT5Dir = Join-Path $PSScriptRoot ".venv_t5" +if (Test-Path $VenvT5Dir) { Remove-Item -Recurse -Force $VenvT5Dir } +New-Item -ItemType Directory -Path $VenvT5Dir -Force | Out-Null +$prevEAP_t5 = $ErrorActionPreference +$ErrorActionPreference = "Continue" +pip install --target $VenvT5Dir --no-deps "transformers==5.2.0" 2>&1 | Out-Null +if ($LASTEXITCODE -ne 0) { + Write-Host "[FAIL] Could not install transformers 5.2.0 into .venv_t5/" -ForegroundColor Red + $ErrorActionPreference = $prevEAP_t5 + exit 1 +} +pip install --target $VenvT5Dir --no-deps "huggingface_hub==1.3.0" 2>&1 | Out-Null +if ($LASTEXITCODE -ne 0) { + Write-Host "[FAIL] Could not install huggingface_hub 1.3.0 into .venv_t5/" -ForegroundColor Red + $ErrorActionPreference = $prevEAP_t5 + exit 1 +} +$ErrorActionPreference = $prevEAP_t5 +Write-Host "[OK] Transformers 5.x pre-installed to .venv_t5/" -ForegroundColor Green + # ========================================================================== # PHASE 3.5: Install OpenSSL dev (for HTTPS support in llama-server) # ========================================================================== diff --git a/setup.sh b/setup.sh index 07ea131125..8e5ae2ac92 100755 --- a/setup.sh +++ b/setup.sh @@ -179,10 +179,24 @@ if [ "$IS_COLAB" = true ]; then else # Local: create venv (always start fresh to preserve correct install order) rm -rf .venv + rm -rf .venv_overlay # Remove legacy overlay (no longer used) + rm -rf .venv_t5 # Will be rebuilt below "$BEST_PY" -m venv .venv source .venv/bin/activate install_python_stack - + + # ── 6b. Pre-install transformers 5.x into .venv_t5/ ── + # Models like GLM-4.7-Flash need transformers>=5.2.0. Instead of pip-installing + # at runtime (slow, ~10-15s), we pre-install into a separate directory. + # The training subprocess just prepends .venv_t5/ to sys.path — instant switch. + echo "" + echo " Pre-installing transformers 5.x for newer model support..." + VENV_T5_DIR="$SCRIPT_DIR/.venv_t5" + mkdir -p "$VENV_T5_DIR" + run_quiet "pip install transformers 5.x" pip install --target "$VENV_T5_DIR" --no-deps "transformers==5.2.0" + run_quiet "pip install huggingface_hub for t5" pip install --target "$VENV_T5_DIR" --no-deps "huggingface_hub==1.3.0" + echo "✅ Transformers 5.x pre-installed to .venv_t5/" + # ── 7. WSL: pre-install GGUF build dependencies ── # On WSL, sudo requires a password and can't be entered during GGUF export # (runs in a non-interactive subprocess). Install build deps here instead. diff --git a/studio/backend/core/__init__.py b/studio/backend/core/__init__.py index a4a700d45b..2b1e3246c0 100644 --- a/studio/backend/core/__init__.py +++ b/studio/backend/core/__init__.py @@ -1,30 +1,18 @@ """ Unified core module for Unsloth backend + +Imports are LAZY (via __getattr__) so that training subprocesses can +import core.training.worker without pulling in heavy ML dependencies +like unsloth, transformers, or torch before the version activation +code has a chance to run. """ -# Inference -from .inference import InferenceBackend, get_inference_backend - -# Training -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 - -# Utilities (from utils) -from utils.paths import normalize_path, is_local_path, is_model_cached -from utils.utils import without_hf_auth, format_error_message -from utils.hardware import get_device, is_apple_silicon, clear_gpu_cache, get_gpu_memory_info, log_gpu_memory, DeviceType -from utils.datasets import format_and_template_dataset - __all__ = [ # Inference 'InferenceBackend', 'get_inference_backend', # Training - 'UnslothTrainer', - 'get_trainer', 'get_training_backend', 'TrainingBackend', 'TrainingProgress', @@ -50,3 +38,72 @@ __all__ = [ 'clear_gpu_cache', 'DeviceType', ] + + +def __getattr__(name): + # Inference + if name in ('InferenceBackend', 'get_inference_backend'): + from .inference import InferenceBackend, get_inference_backend + globals()['InferenceBackend'] = InferenceBackend + globals()['get_inference_backend'] = get_inference_backend + return globals()[name] + + # Training + if name in ('TrainingBackend', 'get_training_backend', 'TrainingProgress'): + from .training import TrainingBackend, get_training_backend, TrainingProgress + globals()['TrainingBackend'] = TrainingBackend + globals()['get_training_backend'] = get_training_backend + globals()['TrainingProgress'] = TrainingProgress + return globals()[name] + + # Config (from utils.models) + if name in ('is_vision_model', 'ModelConfig', 'scan_trained_loras', + 'load_model_defaults', 'get_base_model_from_lora'): + from utils.models import ( + is_vision_model, ModelConfig, scan_trained_loras, + load_model_defaults, get_base_model_from_lora, + ) + globals()['is_vision_model'] = is_vision_model + globals()['ModelConfig'] = ModelConfig + globals()['scan_trained_loras'] = scan_trained_loras + globals()['load_model_defaults'] = load_model_defaults + globals()['get_base_model_from_lora'] = get_base_model_from_lora + return globals()[name] + + # Paths + if name in ('normalize_path', 'is_local_path', 'is_model_cached'): + from utils.paths import normalize_path, is_local_path, is_model_cached + globals()['normalize_path'] = normalize_path + globals()['is_local_path'] = is_local_path + globals()['is_model_cached'] = is_model_cached + return globals()[name] + + # Utils + if name in ('without_hf_auth', 'format_error_message'): + from utils.utils import without_hf_auth, format_error_message + globals()['without_hf_auth'] = without_hf_auth + globals()['format_error_message'] = format_error_message + return globals()[name] + + # Hardware + if name in ('get_device', 'is_apple_silicon', 'clear_gpu_cache', + 'get_gpu_memory_info', 'log_gpu_memory', 'DeviceType'): + from utils.hardware import ( + get_device, is_apple_silicon, clear_gpu_cache, + get_gpu_memory_info, log_gpu_memory, DeviceType, + ) + globals()['get_device'] = get_device + globals()['is_apple_silicon'] = is_apple_silicon + globals()['clear_gpu_cache'] = clear_gpu_cache + globals()['get_gpu_memory_info'] = get_gpu_memory_info + globals()['log_gpu_memory'] = log_gpu_memory + globals()['DeviceType'] = DeviceType + return globals()[name] + + # Datasets + if name == 'format_and_template_dataset': + from utils.datasets import format_and_template_dataset + globals()['format_and_template_dataset'] = format_and_template_dataset + return format_and_template_dataset + + raise AttributeError(f"module 'core' has no attribute {name!r}") diff --git a/studio/backend/core/export/__init__.py b/studio/backend/core/export/__init__.py index 66154f48eb..0a883f5f3f 100644 --- a/studio/backend/core/export/__init__.py +++ b/studio/backend/core/export/__init__.py @@ -1,9 +1,17 @@ """ Export submodule - Model export operations + +The default get_export_backend() returns an ExportOrchestrator that +delegates to a subprocess. The original ExportBackend runs inside +the subprocess and can be imported directly from .export when needed. """ -from .export import ExportBackend, get_export_backend +from .orchestrator import ExportOrchestrator, get_export_backend + +# Expose ExportOrchestrator as ExportBackend for backward compat +ExportBackend = ExportOrchestrator __all__ = [ 'ExportBackend', + 'ExportOrchestrator', 'get_export_backend', ] diff --git a/studio/backend/core/export/export.py b/studio/backend/core/export/export.py index 3d3db2d560..c0f53f2e0e 100644 --- a/studio/backend/core/export/export.py +++ b/studio/backend/core/export/export.py @@ -453,7 +453,14 @@ class ExportBackend: # Write export metadata so the Chat page can identify the base model self._write_export_metadata(abs_save_dir) - logger.info(f"GGUF model saved successfully in {abs_save_dir}") + # Log final file locations (after relocation) so it's clear + # where the GGUF files actually ended up. + final_ggufs = sorted(glob.glob(os.path.join(abs_save_dir, "*.gguf"))) + logger.info( + "GGUF export complete. Final files in %s:\n %s", + abs_save_dir, + "\n ".join(os.path.basename(f) for f in final_ggufs) or "(none)", + ) # Push to hub if requested if push_to_hub: diff --git a/studio/backend/core/export/orchestrator.py b/studio/backend/core/export/orchestrator.py new file mode 100644 index 0000000000..ec18d3533e --- /dev/null +++ b/studio/backend/core/export/orchestrator.py @@ -0,0 +1,392 @@ +""" +Export orchestrator — subprocess-based. + +Provides the same API as ExportBackend, but delegates all ML work +to a persistent subprocess. The subprocess is spawned on first checkpoint +load and stays alive for subsequent export operations. + +When switching between checkpoints that need different transformers versions, +the old subprocess is killed and a new one is spawned with the correct version. + +Pattern follows core/inference/orchestrator.py. +""" +import atexit +import logging +import multiprocessing as mp +import queue +import threading +import time +from pathlib import Path +from typing import Any, List, Optional, Tuple + +logger = logging.getLogger(__name__) + +_CTX = mp.get_context("spawn") + + +class ExportOrchestrator: + """ + Export backend orchestrator — subprocess-based. + + Exposes the same API surface as ExportBackend so routes/export.py + needs minimal changes. Internally, all heavy ML operations happen in + a persistent subprocess. + """ + + def __init__(self): + # Subprocess state + self._proc: Optional[mp.Process] = None + self._cmd_queue: Any = None + self._resp_queue: Any = None + self._lock = threading.Lock() + + # Local state mirrors (updated from subprocess responses) + self.current_checkpoint: Optional[str] = None + self.is_vision: bool = False + self.is_peft: bool = False + + atexit.register(self._cleanup) + logger.info("ExportOrchestrator initialized (subprocess mode)") + + # ------------------------------------------------------------------ + # Subprocess lifecycle + # ------------------------------------------------------------------ + + def _spawn_subprocess(self, config: dict) -> None: + """Spawn a new export subprocess.""" + from .worker import run_export_process + + self._cmd_queue = _CTX.Queue() + self._resp_queue = _CTX.Queue() + + self._proc = _CTX.Process( + target=run_export_process, + kwargs={ + "cmd_queue": self._cmd_queue, + "resp_queue": self._resp_queue, + "config": config, + }, + daemon=True, + ) + self._proc.start() + logger.info("Export subprocess started (pid=%s)", self._proc.pid) + + def _shutdown_subprocess(self, timeout: float = 10.0) -> None: + """Gracefully shut down the export subprocess.""" + if self._proc is None or not self._proc.is_alive(): + self._proc = None + return + + # 1. Drain stale responses + self._drain_queue() + + # 2. Send shutdown command + try: + self._cmd_queue.put({"type": "shutdown"}) + except (OSError, ValueError): + pass + + # 3. Wait for graceful shutdown + try: + self._proc.join(timeout=timeout) + except Exception: + pass + + # 4. Force kill if still alive + if self._proc is not None and self._proc.is_alive(): + logger.warning("Export subprocess did not exit gracefully, terminating") + try: + self._proc.terminate() + self._proc.join(timeout=5) + except Exception: + pass + if self._proc is not None and self._proc.is_alive(): + logger.warning("Subprocess still alive after terminate, killing") + try: + self._proc.kill() + self._proc.join(timeout=3) + except Exception: + pass + + self._proc = None + self._cmd_queue = None + self._resp_queue = None + logger.info("Export subprocess shut down") + + def _cleanup(self): + """atexit handler.""" + self._shutdown_subprocess(timeout=5.0) + + def _ensure_subprocess_alive(self) -> bool: + """Check if subprocess is alive.""" + return self._proc is not None and self._proc.is_alive() + + # ------------------------------------------------------------------ + # Queue helpers + # ------------------------------------------------------------------ + + def _send_cmd(self, cmd: dict) -> None: + """Send a command to the subprocess.""" + if self._cmd_queue is None: + raise RuntimeError("No export subprocess running") + try: + self._cmd_queue.put(cmd) + except (OSError, ValueError) as exc: + raise RuntimeError(f"Failed to send command to subprocess: {exc}") + + def _read_resp(self, timeout: float = 1.0) -> Optional[dict]: + """Read a response from the subprocess (non-blocking with timeout).""" + if self._resp_queue is None: + return None + try: + return self._resp_queue.get(timeout=timeout) + except queue.Empty: + return None + except (EOFError, OSError, ValueError): + return None + + def _wait_response( + self, expected_type: str, timeout: float = 3600.0 + ) -> dict: + """Block until a response of the expected type arrives. + + Export operations can take a very long time — GGUF conversion for + large models (30B+) easily takes 20-30 minutes. Default timeout + is 1 hour. + """ + deadline = time.monotonic() + timeout + + while time.monotonic() < deadline: + remaining = max(0.1, deadline - time.monotonic()) + resp = self._read_resp(timeout=min(remaining, 2.0)) + + if resp is None: + # Check subprocess health + if not self._ensure_subprocess_alive(): + raise RuntimeError("Export subprocess crashed during wait") + continue + + rtype = resp.get("type", "") + + if rtype == expected_type: + return resp + + if rtype == "error": + error_msg = resp.get("error", "Unknown error") + raise RuntimeError(f"Subprocess error: {error_msg}") + + if rtype == "status": + logger.info("Export subprocess status: %s", resp.get("message", "")) + continue + + # Other response types during wait — skip + logger.debug( + "Skipping response type '%s' while waiting for '%s'", + rtype, expected_type, + ) + + raise RuntimeError( + f"Timeout waiting for '{expected_type}' response after {timeout}s" + ) + + def _drain_queue(self) -> list: + """Drain all pending responses.""" + events = [] + if self._resp_queue is None: + return events + while True: + try: + events.append(self._resp_queue.get_nowait()) + except queue.Empty: + return events + except (EOFError, OSError, ValueError): + return events + + # ------------------------------------------------------------------ + # Public API — same interface as ExportBackend + # ------------------------------------------------------------------ + + 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. + + Always spawns a fresh subprocess to ensure a clean Python interpreter. + """ + project_root = str( + Path(__file__).resolve().parent.parent.parent.parent.parent + ) + + sub_config = { + "project_root": project_root, + "checkpoint_path": checkpoint_path, + "max_seq_length": max_seq_length, + "load_in_4bit": load_in_4bit, + } + + # Always kill existing subprocess and spawn fresh. + if self._ensure_subprocess_alive(): + self._shutdown_subprocess() + elif self._proc is not None: + self._shutdown_subprocess(timeout=2) + + logger.info("Spawning fresh export subprocess for '%s'", checkpoint_path) + self._spawn_subprocess(sub_config) + + try: + resp = self._wait_response("loaded", timeout=300) + except RuntimeError as exc: + self._shutdown_subprocess(timeout=5) + self.current_checkpoint = None + self.is_vision = False + self.is_peft = False + return False, str(exc) + + if resp.get("success"): + self.current_checkpoint = resp.get("checkpoint") + self.is_vision = resp.get("is_vision", False) + self.is_peft = resp.get("is_peft", False) + logger.info("Checkpoint '%s' loaded in subprocess", checkpoint_path) + return True, resp.get("message", "Loaded successfully") + else: + error = resp.get("message", "Failed to load checkpoint") + logger.error("Failed to load checkpoint: %s", error) + self.current_checkpoint = None + self.is_vision = False + self.is_peft = False + return False, error + + 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 PEFT model.""" + return self._run_export("merged", { + "save_directory": save_directory, + "format_type": format_type, + "push_to_hub": push_to_hub, + "repo_id": repo_id, + "hf_token": hf_token, + "private": private, + }) + + 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, + base_model_id: Optional[str] = None, + ) -> Tuple[bool, str]: + """Export base model (non-PEFT).""" + return self._run_export("base", { + "save_directory": save_directory, + "push_to_hub": push_to_hub, + "repo_id": repo_id, + "hf_token": hf_token, + "private": private, + "base_model_id": base_model_id, + }) + + 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.""" + return self._run_export("gguf", { + "save_directory": save_directory, + "quantization_method": quantization_method, + "push_to_hub": push_to_hub, + "repo_id": repo_id, + "hf_token": hf_token, + }) + + 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.""" + return self._run_export("lora", { + "save_directory": save_directory, + "push_to_hub": push_to_hub, + "repo_id": repo_id, + "hf_token": hf_token, + "private": private, + }) + + def _run_export(self, export_type: str, params: dict) -> Tuple[bool, str]: + """Send an export command to the subprocess and wait for result.""" + if not self._ensure_subprocess_alive(): + return False, "No export subprocess running. Load a checkpoint first." + + cmd = {"type": "export", "export_type": export_type, **params} + + try: + self._send_cmd(cmd) + resp = self._wait_response( + f"export_{export_type}_done", + timeout=3600, # GGUF for 30B+ models can take 30+ min + ) + return resp.get("success", False), resp.get("message", "") + except RuntimeError as exc: + return False, str(exc) + + def cleanup_memory(self) -> bool: + """Cleanup export-related models from memory.""" + if not self._ensure_subprocess_alive(): + # No subprocess — just clear local state + self.current_checkpoint = None + self.is_vision = False + self.is_peft = False + return True + + try: + self._send_cmd({"type": "cleanup"}) + resp = self._wait_response("cleanup_done", timeout=30) + success = resp.get("success", False) + except RuntimeError: + success = False + + # Shut down subprocess after cleanup — no model loaded + self._shutdown_subprocess() + + self.current_checkpoint = None + self.is_vision = False + self.is_peft = False + return success + + def scan_checkpoints( + self, outputs_dir: str = "./outputs" + ) -> List[Tuple[str, list]]: + """Scan for checkpoints — no ML imports needed, runs locally.""" + from utils.models.checkpoints import scan_checkpoints + return scan_checkpoints(outputs_dir=outputs_dir) + + +# ========== GLOBAL INSTANCE ========== +_export_backend = None + + +def get_export_backend() -> ExportOrchestrator: + """Get global export backend instance (orchestrator).""" + global _export_backend + if _export_backend is None: + _export_backend = ExportOrchestrator() + return _export_backend diff --git a/studio/backend/core/export/worker.py b/studio/backend/core/export/worker.py new file mode 100644 index 0000000000..9e7e72e9dd --- /dev/null +++ b/studio/backend/core/export/worker.py @@ -0,0 +1,350 @@ +""" +Export subprocess entry point. + +Each export session runs in a persistent subprocess (mp.get_context("spawn")). +This gives us a clean Python interpreter with no stale module state — +solving the transformers version-switching problem completely. + +The subprocess stays alive while a model is loaded, accepting commands +(load, export_merged, export_base, export_gguf, export_lora, cleanup, +shutdown) via mp.Queue. + +Pattern follows core/inference/worker.py and core/training/worker.py. +""" +from __future__ import annotations + +import logging +import os +import sys +import time +import traceback +from typing import Any + +logger = logging.getLogger(__name__) + + +def _activate_transformers_version(model_name: str, project_root: str) -> None: + """Activate the correct transformers version BEFORE any ML imports. + + If the model needs transformers 5.x, prepend the pre-installed .venv_t5/ + directory to sys.path. Otherwise do nothing (default 4.57.x in .venv/). + """ + # Ensure backend is on path for utils imports + backend_path = os.path.join(project_root, "studio", "backend") + if backend_path not in sys.path: + sys.path.insert(0, backend_path) + + from utils.transformers_version import needs_transformers_5, _resolve_base_model + + resolved = _resolve_base_model(model_name) + if needs_transformers_5(resolved): + venv_t5 = os.path.join(project_root, ".venv_t5") + if os.path.isdir(venv_t5): + sys.path.insert(0, venv_t5) + logger.info("Activated transformers 5.x from %s", venv_t5) + else: + + # Fallback: pip install at runtime (slower, ~10-15s) + logger.warning(".venv_t5 not found at %s — installing at runtime", venv_t5) + import subprocess as sp + os.makedirs(venv_t5, exist_ok=True) + r1 = sp.run( + [sys.executable, "-m", "pip", "install", "--target", venv_t5, + "--no-deps", "transformers==5.2.0"], + stdout=sp.PIPE, stderr=sp.STDOUT, + ) + r2 = sp.run( + [sys.executable, "-m", "pip", "install", "--target", venv_t5, + "--no-deps", "huggingface_hub==1.3.0"], + stdout=sp.PIPE, stderr=sp.STDOUT, + ) + if r1.returncode != 0 or r2.returncode != 0: + raise RuntimeError( + f"Failed to install transformers 5.x into {venv_t5}. " + f"pip returncode: transformers={r1.returncode}, huggingface_hub={r2.returncode}" + ) + sys.path.insert(0, venv_t5) + # Propagate to child subprocesses (e.g. GGUF converter) + _pp = os.environ.get("PYTHONPATH", "") + os.environ["PYTHONPATH"] = venv_t5 + (os.pathsep + _pp if _pp else "") + else: + logger.info("Using default transformers (4.57.x) for %s", model_name) + + +def _send_response(resp_queue: Any, response: dict) -> None: + """Send a response to the parent process.""" + try: + resp_queue.put(response) + except (OSError, ValueError) as exc: + logger.error("Failed to send response: %s", exc) + + +def _handle_load(backend, cmd: dict, resp_queue: Any) -> None: + """Handle a load_checkpoint command.""" + checkpoint_path = cmd["checkpoint_path"] + max_seq_length = cmd.get("max_seq_length", 2048) + load_in_4bit = cmd.get("load_in_4bit", True) + + try: + _send_response(resp_queue, { + "type": "status", + "message": f"Loading checkpoint: {checkpoint_path}", + "ts": time.time(), + }) + + success, message = backend.load_checkpoint( + checkpoint_path=checkpoint_path, + max_seq_length=max_seq_length, + load_in_4bit=load_in_4bit, + ) + + _send_response(resp_queue, { + "type": "loaded", + "success": success, + "message": message, + "checkpoint": checkpoint_path if success else None, + "is_vision": backend.is_vision if success else False, + "is_peft": backend.is_peft if success else False, + "ts": time.time(), + }) + + except Exception as exc: + _send_response(resp_queue, { + "type": "loaded", + "success": False, + "message": str(exc), + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + + +def _handle_export(backend, cmd: dict, resp_queue: Any) -> None: + """Handle any export command (merged, base, gguf, lora).""" + export_type = cmd["export_type"] # "merged", "base", "gguf", "lora" + response_type = f"export_{export_type}_done" + + try: + if export_type == "merged": + success, message = backend.export_merged_model( + save_directory=cmd.get("save_directory", ""), + format_type=cmd.get("format_type", "16-bit (FP16)"), + push_to_hub=cmd.get("push_to_hub", False), + repo_id=cmd.get("repo_id"), + hf_token=cmd.get("hf_token"), + private=cmd.get("private", False), + ) + elif export_type == "base": + success, message = backend.export_base_model( + save_directory=cmd.get("save_directory", ""), + push_to_hub=cmd.get("push_to_hub", False), + repo_id=cmd.get("repo_id"), + hf_token=cmd.get("hf_token"), + private=cmd.get("private", False), + base_model_id=cmd.get("base_model_id"), + ) + elif export_type == "gguf": + success, message = backend.export_gguf( + save_directory=cmd.get("save_directory", ""), + quantization_method=cmd.get("quantization_method", "Q4_K_M"), + push_to_hub=cmd.get("push_to_hub", False), + repo_id=cmd.get("repo_id"), + hf_token=cmd.get("hf_token"), + ) + elif export_type == "lora": + success, message = backend.export_lora_adapter( + save_directory=cmd.get("save_directory", ""), + push_to_hub=cmd.get("push_to_hub", False), + repo_id=cmd.get("repo_id"), + hf_token=cmd.get("hf_token"), + private=cmd.get("private", False), + ) + else: + success, message = False, f"Unknown export type: {export_type}" + + _send_response(resp_queue, { + "type": response_type, + "success": success, + "message": message, + "ts": time.time(), + }) + + except Exception as exc: + _send_response(resp_queue, { + "type": response_type, + "success": False, + "message": str(exc), + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + + +def _handle_cleanup(backend, resp_queue: Any) -> None: + """Handle a cleanup command.""" + try: + success = backend.cleanup_memory() + _send_response(resp_queue, { + "type": "cleanup_done", + "success": success, + "ts": time.time(), + }) + except Exception as exc: + _send_response(resp_queue, { + "type": "cleanup_done", + "success": False, + "message": str(exc), + "ts": time.time(), + }) + + +def run_export_process( + *, + cmd_queue: Any, + resp_queue: Any, + config: dict, +) -> None: + """Subprocess entrypoint. Persistent — runs command loop until shutdown. + + Args: + cmd_queue: mp.Queue for receiving commands from parent. + resp_queue: mp.Queue for sending responses to parent. + config: Initial configuration dict with checkpoint_path and project_root. + """ + import queue as _queue + + os.environ["TOKENIZERS_PARALLELISM"] = "false" + + project_root = config["project_root"] + checkpoint_path = config["checkpoint_path"] + + # ── 1. Activate correct transformers version BEFORE any ML imports ── + try: + _activate_transformers_version(checkpoint_path, project_root) + except Exception as exc: + _send_response(resp_queue, { + "type": "error", + "error": f"Failed to activate transformers version: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + return + + # ── 1b. On Windows, check Triton availability (must be before import torch) ── + if sys.platform == "win32": + try: + import triton # noqa: F401 + logger.info("Triton available — torch.compile enabled") + except ImportError: + os.environ["TORCHDYNAMO_DISABLE"] = "1" + logger.warning( + "Triton not found on Windows — torch.compile disabled. " + 'Install for better performance: pip install "triton-windows<3.7"' + ) + + # ── 2. Import ML libraries (fresh in this clean process) ── + try: + _send_response(resp_queue, { + "type": "status", + "message": "Importing ML libraries...", + "ts": time.time(), + }) + + backend_path = os.path.join(project_root, "studio", "backend") + if backend_path not in sys.path: + sys.path.insert(0, backend_path) + + from core.export.export import ExportBackend + + import transformers + logger.info("Export subprocess loaded transformers %s", transformers.__version__) + + except Exception as exc: + _send_response(resp_queue, { + "type": "error", + "error": f"Failed to import ML libraries: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + return + + # ── 3. Create export backend and load initial checkpoint ── + try: + backend = ExportBackend() + + _handle_load(backend, config, resp_queue) + + except Exception as exc: + _send_response(resp_queue, { + "type": "error", + "error": f"Failed to initialize export backend: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + return + + # ── 4. Command loop — process commands until shutdown ── + logger.info("Export subprocess ready, entering command loop") + + while True: + try: + cmd = cmd_queue.get(timeout=1.0) + except _queue.Empty: + continue + except (EOFError, OSError): + logger.info("Command queue closed, shutting down") + return + + if cmd is None: + continue + + cmd_type = cmd.get("type", "") + logger.info("Received command: %s", cmd_type) + + try: + if cmd_type == "load": + # Load a new checkpoint (reusing this subprocess) + backend.cleanup_memory() + _handle_load(backend, cmd, resp_queue) + + elif cmd_type == "export": + _handle_export(backend, cmd, resp_queue) + + elif cmd_type == "cleanup": + _handle_cleanup(backend, resp_queue) + + elif cmd_type == "status": + _send_response(resp_queue, { + "type": "status_response", + "checkpoint": backend.current_checkpoint, + "is_vision": backend.is_vision, + "is_peft": backend.is_peft, + "ts": time.time(), + }) + + elif cmd_type == "shutdown": + logger.info("Shutdown command received, cleaning up and exiting") + try: + backend.cleanup_memory() + except Exception: + pass + _send_response(resp_queue, { + "type": "shutdown_ack", + "ts": time.time(), + }) + return + + else: + logger.warning("Unknown command type: %s", cmd_type) + _send_response(resp_queue, { + "type": "error", + "error": f"Unknown command type: {cmd_type}", + "ts": time.time(), + }) + + except Exception as exc: + logger.error("Error handling command '%s': %s", cmd_type, exc, exc_info=True) + _send_response(resp_queue, { + "type": "error", + "error": f"Command '{cmd_type}' failed: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) diff --git a/studio/backend/core/inference/__init__.py b/studio/backend/core/inference/__init__.py index ff8b75d36a..6d44742d9c 100644 --- a/studio/backend/core/inference/__init__.py +++ b/studio/backend/core/inference/__init__.py @@ -1,11 +1,19 @@ """ Inference submodule - Inference backend for model loading and generation + +The default get_inference_backend() returns an InferenceOrchestrator that +delegates to a subprocess. The original InferenceBackend runs inside +the subprocess and can be imported directly from .inference when needed. """ -from .inference import InferenceBackend, get_inference_backend +from .orchestrator import InferenceOrchestrator, get_inference_backend from .llama_cpp import LlamaCppBackend +# Expose InferenceOrchestrator as InferenceBackend for backward compat +InferenceBackend = InferenceOrchestrator + __all__ = [ 'InferenceBackend', + 'InferenceOrchestrator', 'get_inference_backend', 'LlamaCppBackend', ] diff --git a/studio/backend/core/inference/orchestrator.py b/studio/backend/core/inference/orchestrator.py new file mode 100644 index 0000000000..5db7e56357 --- /dev/null +++ b/studio/backend/core/inference/orchestrator.py @@ -0,0 +1,600 @@ +""" +Inference orchestrator — subprocess-based. + +Provides the same API as InferenceBackend, but delegates all ML work +to a persistent subprocess. The subprocess is spawned on first model load +and stays alive for subsequent requests. + +When switching between models that need different transformers versions +(e.g. GLM-4.7-Flash needs 5.x, Qwen needs 4.57.x), the old subprocess +is killed and a new one is spawned with the correct version. + +Pattern follows core/training/training.py. +""" +import atexit +import base64 +import logging +import multiprocessing as mp +import queue +import threading +import time +import uuid +from io import BytesIO +from pathlib import Path +from typing import Any, Generator, Optional, Tuple, Union + +logger = logging.getLogger(__name__) + +_CTX = mp.get_context("spawn") + + +class InferenceOrchestrator: + """ + Inference backend orchestrator — subprocess-based. + + Exposes the same API surface as InferenceBackend so routes/inference.py + needs minimal changes. Internally, all heavy ML operations happen in + a persistent subprocess. + """ + + def __init__(self): + # Subprocess state + self._proc: Optional[mp.Process] = None + self._cmd_queue: Any = None + self._resp_queue: Any = None + self._cancel_event: Any = None # mp.Event — set to cancel generation instantly + self._lock = threading.Lock() + self._gen_lock = threading.Lock() # Serializes generation — one request at a time + + # Local state mirrors (updated from subprocess responses) + self.active_model_name: Optional[str] = None + self.models: dict = {} + self.loading_models: set = set() + self.loaded_local_models: list = [] + 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", + ] + + # Version tracking for subprocess reuse + self._current_transformers_major: Optional[str] = None # "4" or "5" + + atexit.register(self._cleanup) + logger.info("InferenceOrchestrator initialized (subprocess mode)") + + # ------------------------------------------------------------------ + # Subprocess lifecycle + # ------------------------------------------------------------------ + + def _spawn_subprocess(self, config: dict) -> None: + """Spawn a new inference subprocess.""" + from .worker import run_inference_process + + self._cmd_queue = _CTX.Queue() + self._resp_queue = _CTX.Queue() + self._cancel_event = _CTX.Event() + + self._proc = _CTX.Process( + target=run_inference_process, + kwargs={ + "cmd_queue": self._cmd_queue, + "resp_queue": self._resp_queue, + "cancel_event": self._cancel_event, + "config": config, + }, + daemon=True, + ) + self._proc.start() + logger.info("Inference subprocess started (pid=%s)", self._proc.pid) + + def _cancel_generation(self) -> None: + """Cancel any ongoing generation in the subprocess (instant).""" + if self._cancel_event is not None: + self._cancel_event.set() + + def _shutdown_subprocess(self, timeout: float = 10.0) -> None: + """Gracefully shut down the inference subprocess.""" + if self._proc is None or not self._proc.is_alive(): + self._proc = None + return + + # 1. Cancel any ongoing generation first (instant via mp.Event) + self._cancel_generation() + time.sleep(0.5) # Brief wait for generation to stop + + # 2. Drain stale responses from queue + self._drain_queue() + + # 3. Send shutdown command + try: + self._cmd_queue.put({"type": "shutdown"}) + except (OSError, ValueError): + pass + + # 4. Wait for graceful shutdown + try: + self._proc.join(timeout=timeout) + except Exception: + pass + + # 5. Force kill if still alive + if self._proc is not None and self._proc.is_alive(): + logger.warning("Inference subprocess did not exit gracefully, terminating") + try: + self._proc.terminate() + self._proc.join(timeout=5) + except Exception: + pass + if self._proc is not None and self._proc.is_alive(): + logger.warning("Subprocess still alive after terminate, killing") + try: + self._proc.kill() + self._proc.join(timeout=3) + except Exception: + pass + + self._proc = None + self._cmd_queue = None + self._resp_queue = None + self._cancel_event = None + logger.info("Inference subprocess shut down") + + def _cleanup(self): + """atexit handler.""" + self._shutdown_subprocess(timeout=5.0) + + def _ensure_subprocess_alive(self) -> bool: + """Check if subprocess is alive.""" + return self._proc is not None and self._proc.is_alive() + + # ------------------------------------------------------------------ + # Queue helpers + # ------------------------------------------------------------------ + + def _send_cmd(self, cmd: dict) -> None: + """Send a command to the subprocess.""" + if self._cmd_queue is None: + raise RuntimeError("No inference subprocess running") + try: + self._cmd_queue.put(cmd) + except (OSError, ValueError) as exc: + raise RuntimeError(f"Failed to send command to subprocess: {exc}") + + def _read_resp(self, timeout: float = 1.0) -> Optional[dict]: + """Read a response from the subprocess (non-blocking with timeout).""" + if self._resp_queue is None: + return None + try: + return self._resp_queue.get(timeout=timeout) + except queue.Empty: + return None + except (EOFError, OSError, ValueError): + return None + + def _wait_response( + self, expected_type: str, timeout: float = 120.0 + ) -> dict: + """Block until a response of the expected type arrives. + + Also handles 'status' and 'error' events during the wait. + Returns the matching response dict. + Raises RuntimeError on timeout or subprocess crash. + """ + deadline = time.monotonic() + timeout + + while time.monotonic() < deadline: + remaining = max(0.1, deadline - time.monotonic()) + resp = self._read_resp(timeout=min(remaining, 1.0)) + + if resp is None: + # Check subprocess health + if not self._ensure_subprocess_alive(): + raise RuntimeError("Inference subprocess crashed during wait") + continue + + rtype = resp.get("type", "") + + if rtype == expected_type: + return resp + + if rtype == "error": + error_msg = resp.get("error", "Unknown error") + raise RuntimeError(f"Subprocess error: {error_msg}") + + if rtype == "status": + logger.info("Subprocess status: %s", resp.get("message", "")) + continue + + # Other response types during wait — skip + logger.debug("Skipping response type '%s' while waiting for '%s'", rtype, expected_type) + + raise RuntimeError( + f"Timeout waiting for '{expected_type}' response after {timeout}s" + ) + + def _drain_queue(self) -> list: + """Drain all pending responses.""" + events = [] + if self._resp_queue is None: + return events + while True: + try: + events.append(self._resp_queue.get_nowait()) + except queue.Empty: + return events + except (EOFError, OSError, ValueError): + return events + + def _drain_until_gen_done(self, timeout: float = 5.0) -> None: + """Consume resp_queue events until gen_done/gen_error, discarding them. + + Called after cancel to ensure stale tokens from the cancelled + generation don't leak into the next request. + """ + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + resp = self._read_resp(timeout=min(0.5, deadline - time.monotonic())) + if resp is None: + if not self._ensure_subprocess_alive(): + return + continue + rtype = resp.get("type", "") + if rtype in ("gen_done", "gen_error"): + return + logger.warning("Timed out waiting for gen_done after cancel") + + # ------------------------------------------------------------------ + # Public API — same interface as InferenceBackend + # ------------------------------------------------------------------ + + 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 a model for inference. + + Always spawns a fresh subprocess for each model load. This ensures + a clean Python interpreter — no stale unsloth patches, torch.compile + caches, or inspect.getsource() failures from a previous model. + """ + from utils.transformers_version import needs_transformers_5 + + model_name = config.identifier + self.loading_models.add(model_name) + + try: + needed_major = "5" if needs_transformers_5(model_name) else "4" + project_root = str(Path(__file__).resolve().parent.parent.parent.parent.parent) + + # Build config dict for subprocess + sub_config = { + "project_root": project_root, + "model_name": model_name, + "max_seq_length": max_seq_length, + "load_in_4bit": load_in_4bit, + "hf_token": hf_token or "", + "gguf_variant": getattr(config, "gguf_variant", None), + } + + # Always kill existing subprocess and spawn fresh. + # Reusing a subprocess after unsloth patches torch internals + # causes inspect.getsource() failures on the next model load. + if self._ensure_subprocess_alive(): + self._cancel_generation() + time.sleep(0.3) + self._shutdown_subprocess() + + elif self._proc is not None: + # Dead subprocess — clean up + self._shutdown_subprocess(timeout=2) + + logger.info( + "Spawning fresh inference subprocess for '%s' (transformers %s.x)", + model_name, needed_major, + ) + self._spawn_subprocess(sub_config) + resp = self._wait_response("loaded", timeout=180) + + # Update local state from response + if resp.get("success"): + self._current_transformers_major = needed_major + model_info = resp.get("model_info", {}) + self.active_model_name = model_info.get("identifier", model_name) + self.models[self.active_model_name] = { + "is_vision": model_info.get("is_vision", False), + "is_lora": model_info.get("is_lora", False), + "display_name": model_info.get("display_name", model_name), + } + self.loading_models.discard(model_name) + logger.info("Model '%s' loaded successfully in subprocess", model_name) + return True + else: + error = resp.get("error", "Failed to load model") + self.loading_models.discard(model_name) + self.active_model_name = None + self.models.clear() + raise Exception(error) + + except Exception: + self.loading_models.discard(model_name) + self.active_model_name = None + self.models.clear() + raise + + def unload_model(self, model_name: str) -> bool: + """Unload a model from the subprocess.""" + if not self._ensure_subprocess_alive(): + # No subprocess — just clear local state + self.models.pop(model_name, None) + if self.active_model_name == model_name: + self.active_model_name = None + return True + + try: + self._send_cmd({ + "type": "unload", + "model_name": model_name, + }) + resp = self._wait_response("unloaded", timeout=30) + + # Update local state + self.models.pop(model_name, None) + if self.active_model_name == model_name: + self.active_model_name = None + + logger.info("Model '%s' unloaded from subprocess", model_name) + return True + + except Exception as exc: + logger.error("Error unloading model '%s': %s", model_name, exc) + # Clear local state anyway + self.models.pop(model_name, None) + if self.active_model_name == model_name: + self.active_model_name = None + return False + + def generate_chat_response( + self, + messages: list, + system_prompt: str = "", + image=None, + temperature: float = 0.7, + top_p: float = 0.9, + top_k: int = 40, + min_p: float = 0.0, + max_new_tokens: int = 256, + repetition_penalty: float = 1.1, + cancel_event=None, + ) -> Generator[str, None, None]: + """Generate response, streaming tokens from subprocess.""" + yield from self._generate_inner( + messages=messages, + system_prompt=system_prompt, + image=image, + temperature=temperature, + top_p=top_p, + top_k=top_k, + min_p=min_p, + max_new_tokens=max_new_tokens, + repetition_penalty=repetition_penalty, + cancel_event=cancel_event, + use_adapter=None, + ) + + def generate_with_adapter_control( + self, + use_adapter: Optional[Union[bool, str]] = None, + cancel_event=None, + **gen_kwargs, + ) -> Generator[str, None, None]: + """Generate with adapter control, streaming tokens from subprocess.""" + yield from self._generate_inner( + use_adapter=use_adapter, + cancel_event=cancel_event, + **gen_kwargs, + ) + + def _generate_inner( + self, + messages: list = None, + system_prompt: str = "", + image=None, + temperature: float = 0.7, + top_p: float = 0.9, + top_k: int = 40, + min_p: float = 0.0, + max_new_tokens: int = 256, + repetition_penalty: float = 1.1, + cancel_event=None, + use_adapter=None, + ) -> Generator[str, None, None]: + """Inner generation logic — sends command to subprocess, yields tokens. + + Serialized by _gen_lock: only one generation runs at a time. + This prevents concurrent readers from consuming each other's + tokens off the shared resp_queue. + """ + if not self._ensure_subprocess_alive(): + yield "Error: Inference subprocess is not running" + return + + if not self.active_model_name: + yield "Error: No active model" + return + + # Serialize generation — single GPU, one generation at a time. + # Without this lock, two concurrent readers on the same resp_queue + # can consume and drop each other's token events. + with self._gen_lock: + yield from self._generate_locked( + messages=messages, + system_prompt=system_prompt, + image=image, + temperature=temperature, + top_p=top_p, + top_k=top_k, + min_p=min_p, + max_new_tokens=max_new_tokens, + repetition_penalty=repetition_penalty, + cancel_event=cancel_event, + use_adapter=use_adapter, + ) + + def _generate_locked( + self, + messages: list = None, + system_prompt: str = "", + image=None, + temperature: float = 0.7, + top_p: float = 0.9, + top_k: int = 40, + min_p: float = 0.0, + max_new_tokens: int = 256, + repetition_penalty: float = 1.1, + cancel_event=None, + use_adapter=None, + ) -> Generator[str, None, None]: + """Actual generation logic — must be called under _gen_lock.""" + request_id = str(uuid.uuid4()) + + # Convert PIL Image to base64 if needed + image_b64 = None + if image is not None: + image_b64 = self._pil_to_base64(image) + + cmd = { + "type": "generate", + "request_id": request_id, + "messages": messages or [], + "system_prompt": system_prompt, + "image_base64": image_b64, + "temperature": temperature, + "top_p": top_p, + "top_k": top_k, + "min_p": min_p, + "max_new_tokens": max_new_tokens, + "repetition_penalty": repetition_penalty, + } + + if use_adapter is not None: + cmd["use_adapter"] = use_adapter + + try: + self._send_cmd(cmd) + except RuntimeError as exc: + yield f"Error: {exc}" + return + + # Yield tokens from response queue — we are the only reader + # because _gen_lock is held. + while True: + resp = self._read_resp(timeout=30.0) + + if resp is None: + # Check subprocess health + if not self._ensure_subprocess_alive(): + yield "Error: Inference subprocess crashed during generation" + return + continue + + rtype = resp.get("type", "") + + # Status messages — skip + if rtype == "status": + continue + + # Error without request_id = subprocess-level error + resp_rid = resp.get("request_id") + if rtype == "error" and not resp_rid: + yield f"Error: {resp.get('error', 'Unknown error')}" + return + + if rtype == "token": + # Check cancel from route (e.g. SSE connection closed) + if cancel_event is not None and cancel_event.is_set(): + self._cancel_generation() + # Wait for the subprocess to acknowledge cancellation + # (gen_done/gen_error) so stale events don't leak into + # the next generation request. + self._drain_until_gen_done(timeout=5.0) + return + yield resp.get("text", "") + + elif rtype == "gen_done": + return + + elif rtype == "gen_error": + yield f"Error: {resp.get('error', 'Unknown error')}" + return + + def reset_generation_state(self): + """Cancel any ongoing generation and reset state.""" + self._cancel_generation() + if not self._ensure_subprocess_alive(): + return + try: + self._send_cmd({"type": "reset"}) + except RuntimeError: + pass + + # ------------------------------------------------------------------ + # Local helpers (no subprocess needed) + # ------------------------------------------------------------------ + + def resize_image(self, img, max_size: int = 800): + """Resize image while maintaining aspect ratio. + No ML imports needed — runs locally in parent process. + """ + 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 + + @staticmethod + def _pil_to_base64(img) -> str: + """Convert a PIL Image to base64 string for IPC.""" + buf = BytesIO() + img.save(buf, format="PNG") + return base64.b64encode(buf.getvalue()).decode("ascii") + + 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 check_vision_model_compatibility(self) -> bool: + """Check if current model supports vision.""" + if self.active_model_name and self.active_model_name in self.models: + return self.models[self.active_model_name].get("is_vision", False) + return False + + +# ========== GLOBAL INSTANCE ========== +_inference_backend = None + + +def get_inference_backend() -> InferenceOrchestrator: + """Get global inference backend instance (orchestrator).""" + global _inference_backend + if _inference_backend is None: + _inference_backend = InferenceOrchestrator() + return _inference_backend diff --git a/studio/backend/core/inference/worker.py b/studio/backend/core/inference/worker.py new file mode 100644 index 0000000000..7683d5e314 --- /dev/null +++ b/studio/backend/core/inference/worker.py @@ -0,0 +1,478 @@ +""" +Inference subprocess entry point. + +Each inference session runs in a persistent subprocess (mp.get_context("spawn")). +This gives us a clean Python interpreter with no stale module state — +solving the transformers version-switching problem completely. + +The subprocess stays alive while a model is loaded, accepting commands +(generate, load, unload) via mp.Queue. It exits on shutdown or unload. + +Pattern follows core/training/worker.py. +""" +from __future__ import annotations + +import base64 +import logging +import os +import queue as _queue +import sys +import time +import traceback +from io import BytesIO +from typing import Any + +logger = logging.getLogger(__name__) + + +def _activate_transformers_version(model_name: str, project_root: str) -> None: + """Activate the correct transformers version BEFORE any ML imports. + + If the model needs transformers 5.x, prepend the pre-installed .venv_t5/ + directory to sys.path. Otherwise do nothing (default 4.57.x in .venv/). + """ + # Ensure backend is on path for utils imports + backend_path = os.path.join(project_root, "studio", "backend") + if backend_path not in sys.path: + sys.path.insert(0, backend_path) + + from utils.transformers_version import needs_transformers_5, _resolve_base_model + + resolved = _resolve_base_model(model_name) + if needs_transformers_5(resolved): + venv_t5 = os.path.join(project_root, ".venv_t5") + if os.path.isdir(venv_t5): + sys.path.insert(0, venv_t5) + logger.info("Activated transformers 5.x from %s", venv_t5) + else: + # Fallback: pip install at runtime (slower, ~10-15s) + logger.warning(".venv_t5 not found at %s — installing at runtime", venv_t5) + import subprocess as sp + os.makedirs(venv_t5, exist_ok=True) + r1 = sp.run( + [sys.executable, "-m", "pip", "install", "--target", venv_t5, + "--no-deps", "transformers==5.2.0"], + stdout=sp.PIPE, stderr=sp.STDOUT, + ) + r2 = sp.run( + [sys.executable, "-m", "pip", "install", "--target", venv_t5, + "--no-deps", "huggingface_hub==1.3.0"], + stdout=sp.PIPE, stderr=sp.STDOUT, + ) + if r1.returncode != 0 or r2.returncode != 0: + raise RuntimeError( + f"Failed to install transformers 5.x into {venv_t5}. " + f"pip returncode: transformers={r1.returncode}, huggingface_hub={r2.returncode}" + ) + sys.path.insert(0, venv_t5) + # Propagate to child subprocesses (e.g. GGUF converter) + _pp = os.environ.get("PYTHONPATH", "") + os.environ["PYTHONPATH"] = venv_t5 + (os.pathsep + _pp if _pp else "") + else: + logger.info("Using default transformers (4.57.x) for %s", model_name) + + +def _decode_image(image_base64: str): + """Decode base64 string to PIL.Image.""" + from PIL import Image + image_data = base64.b64decode(image_base64) + return Image.open(BytesIO(image_data)) + + +def _resize_image(img, max_size: int = 800): + """Resize image while maintaining aspect ratio.""" + 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 _send_response(resp_queue: Any, response: dict) -> None: + """Send a response to the parent process.""" + try: + resp_queue.put(response) + except (OSError, ValueError) as exc: + logger.error("Failed to send response: %s", exc) + + +def _build_model_config(config: dict): + """Build a ModelConfig from the config dict.""" + from utils.models import ModelConfig + + model_name = config["model_name"] + hf_token = config.get("hf_token") + hf_token = hf_token if hf_token and hf_token.strip() else None + gguf_variant = config.get("gguf_variant") + + mc = ModelConfig.from_identifier( + model_id=model_name, + hf_token=hf_token, + gguf_variant=gguf_variant, + ) + if not mc: + raise ValueError(f"Invalid model identifier: {model_name}") + return mc + + +def _handle_load(backend, config: dict, resp_queue: Any) -> None: + """Handle a load command: load a model into the backend.""" + try: + mc = _build_model_config(config) + + hf_token = config.get("hf_token") + hf_token = hf_token if hf_token and hf_token.strip() else None + + # Auto-detect quantization for LoRA adapters + load_in_4bit = config.get("load_in_4bit", True) + if mc.is_lora and mc.path: + import json + from pathlib import Path + adapter_cfg_path = Path(mc.path) / "adapter_config.json" + if adapter_cfg_path.exists(): + try: + with open(adapter_cfg_path) as f: + adapter_cfg = json.load(f) + training_method = adapter_cfg.get("unsloth_training_method") + if training_method == "lora" and load_in_4bit: + logger.info("adapter_config.json says lora — setting load_in_4bit=False") + load_in_4bit = False + elif training_method == "qlora" and not load_in_4bit: + logger.info("adapter_config.json says qlora — setting load_in_4bit=True") + load_in_4bit = True + elif not training_method: + if mc.base_model and "-bnb-4bit" not in mc.base_model.lower() and load_in_4bit: + logger.info("No training method, base model has no -bnb-4bit — setting load_in_4bit=False") + load_in_4bit = False + except Exception as e: + logger.warning("Could not read adapter_config.json: %s", e) + + success = backend.load_model( + config=mc, + max_seq_length=config.get("max_seq_length", 2048), + load_in_4bit=load_in_4bit, + hf_token=hf_token, + ) + + if success: + # Build model_info for the parent to mirror + model_info = { + "identifier": mc.identifier, + "display_name": mc.display_name, + "is_vision": mc.is_vision, + "is_lora": mc.is_lora, + "is_gguf": False, + } + _send_response(resp_queue, { + "type": "loaded", + "success": True, + "model_info": model_info, + "ts": time.time(), + }) + else: + _send_response(resp_queue, { + "type": "loaded", + "success": False, + "error": "Failed to load model", + "ts": time.time(), + }) + + except Exception as exc: + _send_response(resp_queue, { + "type": "loaded", + "success": False, + "error": str(exc), + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + + +def _handle_generate( + backend, + cmd: dict, + resp_queue: Any, + cancel_event, +) -> None: + """Handle a generate command: stream tokens back via resp_queue. + + cancel_event is an mp.Event shared with the parent process. + The parent can set it at any time (e.g. user stops generation, + or user loads a new model while generating) and generation + stops within 1-2 tokens. + """ + request_id = cmd.get("request_id", "") + + try: + # Decode image if provided + image = None + image_b64 = cmd.get("image_base64") + if image_b64: + image = _decode_image(image_b64) + image = _resize_image(image) + + # Build generation kwargs + gen_kwargs = { + "messages": cmd["messages"], + "system_prompt": cmd.get("system_prompt", ""), + "image": image, + "temperature": cmd.get("temperature", 0.7), + "top_p": cmd.get("top_p", 0.9), + "top_k": cmd.get("top_k", 40), + "min_p": cmd.get("min_p", 0.0), + "max_new_tokens": cmd.get("max_new_tokens", 256), + "repetition_penalty": cmd.get("repetition_penalty", 1.1), + "cancel_event": cancel_event, + } + + # Choose generation path + use_adapter = cmd.get("use_adapter") + if use_adapter is not None: + generator = backend.generate_with_adapter_control( + use_adapter=use_adapter, + **gen_kwargs, + ) + else: + generator = backend.generate_chat_response(**gen_kwargs) + + for cumulative_text in generator: + # cancel_event is an mp.Event — checked instantly, no queue polling + if cancel_event.is_set(): + logger.info("Generation cancelled for request %s", request_id) + break + + _send_response(resp_queue, { + "type": "token", + "request_id": request_id, + "text": cumulative_text, + "ts": time.time(), + }) + + _send_response(resp_queue, { + "type": "gen_done", + "request_id": request_id, + "ts": time.time(), + }) + + except Exception as exc: + logger.error("Generation error: %s", exc, exc_info=True) + _send_response(resp_queue, { + "type": "gen_error", + "request_id": request_id, + "error": str(exc), + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + + +def _handle_unload(backend, cmd: dict, resp_queue: Any) -> None: + """Handle an unload command.""" + model_name = cmd.get("model_name", "") + try: + if model_name and model_name in backend.models: + backend.unload_model(model_name) + elif backend.active_model_name: + backend.unload_model(backend.active_model_name) + + _send_response(resp_queue, { + "type": "unloaded", + "model_name": model_name, + "ts": time.time(), + }) + except Exception as exc: + logger.error("Unload error: %s", exc) + _send_response(resp_queue, { + "type": "unloaded", + "model_name": model_name, + "error": str(exc), + "ts": time.time(), + }) + + +def run_inference_process( + *, + cmd_queue: Any, + resp_queue: Any, + cancel_event, + config: dict, +) -> None: + """Subprocess entrypoint. Persistent — runs command loop until shutdown. + + Args: + cmd_queue: mp.Queue for receiving commands from parent. + resp_queue: mp.Queue for sending responses to parent. + cancel_event: mp.Event shared with parent — set by parent to cancel generation. + config: Initial configuration dict with model info and project_root. + """ + os.environ["TOKENIZERS_PARALLELISM"] = "false" + + project_root = config["project_root"] + model_name = config["model_name"] + + # ── 1. Activate correct transformers version BEFORE any ML imports ── + try: + _activate_transformers_version(model_name, project_root) + except Exception as exc: + _send_response(resp_queue, { + "type": "error", + "error": f"Failed to activate transformers version: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + return + + # ── 1b. On Windows, check Triton availability (must be before import torch) ── + if sys.platform == "win32": + try: + import triton # noqa: F401 + logger.info("Triton available — torch.compile enabled") + except ImportError: + os.environ["TORCHDYNAMO_DISABLE"] = "1" + logger.warning( + "Triton not found on Windows — torch.compile disabled. " + 'Install for better performance: pip install "triton-windows<3.7"' + ) + + # ── 2. Import ML libraries (fresh in this clean process) ── + try: + _send_response(resp_queue, { + "type": "status", + "message": "Importing ML libraries...", + "ts": time.time(), + }) + + backend_path = os.path.join(project_root, "studio", "backend") + if backend_path not in sys.path: + sys.path.insert(0, backend_path) + + from core.inference.inference import InferenceBackend + + import transformers + logger.info("Subprocess loaded transformers %s", transformers.__version__) + + except Exception as exc: + _send_response(resp_queue, { + "type": "error", + "error": f"Failed to import ML libraries: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + return + + # ── 3. Create inference backend and load initial model ── + try: + backend = InferenceBackend() + + _send_response(resp_queue, { + "type": "status", + "message": "Loading model...", + "ts": time.time(), + }) + + _handle_load(backend, config, resp_queue) + + except Exception as exc: + _send_response(resp_queue, { + "type": "error", + "error": f"Failed to initialize inference backend: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + return + + # ── 4. Command loop — process commands until shutdown ── + # cancel_event is an mp.Event shared with parent — parent can set it + # at any time to cancel generation instantly (no queue polling needed). + logger.info("Inference subprocess ready, entering command loop") + + while True: + try: + cmd = cmd_queue.get(timeout=1.0) + except _queue.Empty: + continue + except (EOFError, OSError): + logger.info("Command queue closed, shutting down") + return + + if cmd is None: + continue + + cmd_type = cmd.get("type", "") + logger.info("Received command: %s", cmd_type) + + try: + if cmd_type == "generate": + cancel_event.clear() + _handle_generate(backend, cmd, resp_queue, cancel_event) + + elif cmd_type == "load": + # Load a new model (reusing this subprocess) + # First unload current model + if backend.active_model_name: + backend.unload_model(backend.active_model_name) + _handle_load(backend, cmd, resp_queue) + + elif cmd_type == "unload": + _handle_unload(backend, cmd, resp_queue) + + elif cmd_type == "cancel": + # Redundant with mp.Event but handle gracefully + cancel_event.set() + logger.info("Cancel command received") + + elif cmd_type == "reset": + cancel_event.set() + backend.reset_generation_state() + _send_response(resp_queue, { + "type": "reset_ack", + "ts": time.time(), + }) + + elif cmd_type == "status": + # Return current status + _send_response(resp_queue, { + "type": "status_response", + "active_model": backend.active_model_name, + "models": { + name: { + "is_vision": info.get("is_vision", False), + "is_lora": info.get("is_lora", False), + } + for name, info in backend.models.items() + }, + "loading": list(backend.loading_models), + "ts": time.time(), + }) + + elif cmd_type == "shutdown": + logger.info("Shutdown command received, exiting") + # Unload all models + for model_name in list(backend.models.keys()): + try: + backend.unload_model(model_name) + except Exception: + pass + _send_response(resp_queue, { + "type": "shutdown_ack", + "ts": time.time(), + }) + return + + else: + logger.warning("Unknown command type: %s", cmd_type) + _send_response(resp_queue, { + "type": "error", + "error": f"Unknown command type: {cmd_type}", + "ts": time.time(), + }) + + except Exception as exc: + logger.error("Error handling command '%s': %s", cmd_type, exc, exc_info=True) + _send_response(resp_queue, { + "type": "error", + "error": f"Command '{cmd_type}' failed: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) diff --git a/studio/backend/core/training/__init__.py b/studio/backend/core/training/__init__.py index 8fc2a6c721..0e0299caf6 100644 --- a/studio/backend/core/training/__init__.py +++ b/studio/backend/core/training/__init__.py @@ -1,12 +1,9 @@ """ Training submodule - Training backends and trainer classes """ -from .trainer import UnslothTrainer, get_trainer, TrainingProgress -from .training import TrainingBackend, get_training_backend +from .training import TrainingBackend, TrainingProgress, get_training_backend __all__ = [ - 'UnslothTrainer', - 'get_trainer', 'TrainingProgress', 'TrainingBackend', 'get_training_backend', diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index f2e43f76f4..aa4b551dd4 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -3,6 +3,7 @@ Unsloth Training Backend Integrates Unsloth training capabilities with the FastAPI backend """ import os +import sys # Prevent tokenizer parallelism deadlocks when datasets uses multiprocessing fork os.environ["TOKENIZERS_PARALLELISM"] = "false" @@ -68,6 +69,7 @@ class UnslothTrainer: self.is_training = False self.should_stop = False self.save_on_stop = True + self.load_in_4bit = True # Track quantization mode for metadata # Model state tracking self.is_vlm = False @@ -114,6 +116,7 @@ class UnslothTrainer: hf_token: Optional[str] = None, is_dataset_multimodal: bool = False) -> bool: """Load model for training (supports both text and vision models)""" + self.load_in_4bit = load_in_4bit # Store for training_meta.json try: if self.model is not None: del self.model @@ -750,6 +753,13 @@ class UnslothTrainer: "include_num_input_tokens_seen": True, # Enable token counting "dataset_num_proc": safe_num_proc(max(1, os.cpu_count() // 4)), } + + # On Windows with transformers 5.x, disable DataLoader multiprocessing + # to avoid issues with modified sys.path (.venv_t5) in spawned workers. + if sys.platform == "win32": + import transformers as _tf + if _tf.__version__.startswith("5."): + config_args["dataloader_num_workers"] = 0 # Add warmup parameter - use warmup_ratio if provided, otherwise warmup_steps if warmup_ratio_val is not None: @@ -1068,6 +1078,7 @@ class UnslothTrainer: # Stopped by user — save model at current checkpoint self.trainer.save_model() self.tokenizer.save_pretrained(output_dir) + self._patch_adapter_config(output_dir) print(f"\nTraining stopped. Model saved to {output_dir}\n") self._update_progress( is_training=False, @@ -1084,6 +1095,7 @@ class UnslothTrainer: # Normal completion self.trainer.save_model() self.tokenizer.save_pretrained(output_dir) + self._patch_adapter_config(output_dir) print(f"\nTraining completed! Model saved to {output_dir}\n") self._update_progress( is_training=False, @@ -1098,6 +1110,36 @@ class UnslothTrainer: finally: self.is_training = False + def _patch_adapter_config(self, output_dir: str) -> None: + """Patch adapter_config.json with unsloth_training_method. + + Values: 'qlora', 'lora', 'FT', 'CPT', 'DPO', 'GRPO', etc. + For LoRA/QLoRA, the distinction comes from load_in_4bit. + """ + config_path = os.path.join(output_dir, "adapter_config.json") + if not os.path.exists(config_path): + logger.info("No adapter_config.json found — skipping training method patch") + return + + try: + with open(config_path, "r") as f: + config = json.load(f) + + # Determine the training method + if self.load_in_4bit: + method = "qlora" + else: + method = "lora" + + config["unsloth_training_method"] = method + logger.info(f"Patching adapter_config.json with unsloth_training_method='{method}'") + + with open(config_path, "w") as f: + json.dump(config, f, indent=2) + + except Exception as e: + logger.warning(f"Failed to patch adapter_config.json: {e}") + def stop_training(self, save: bool = True): """Stop ongoing training""" print(f"\nStopping training (save={save})...") diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 153f4335e3..642b1e8e27 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -1,620 +1,534 @@ """ -Training backend for FastAPI integration -""" -import matplotlib.pyplot as plt -from typing import Any, Generator, Tuple -import logging -import math +Training backend — subprocess orchestrator. -from .trainer import get_trainer, TrainingProgress -from utils.hardware import clear_gpu_cache +Each training job runs in a fresh subprocess (mp.get_context("spawn")), +solving the transformers version-switching problem. The old in-process +UnslothTrainer singleton is only used inside the subprocess (worker.py). + +This file orchestrates the subprocess lifecycle, pumps events from the +worker's mp.Queue, and exposes the same API surface to routes/training.py. + +Pattern follows core/data_recipe/jobs/manager.py. +""" +import math +import multiprocessing as mp +import queue +import threading +import time +import logging +from dataclasses import dataclass, field +from pathlib import Path +from typing import Optional, Tuple, Any + +import matplotlib.pyplot as plt logger = logging.getLogger(__name__) +_CTX = mp.get_context("spawn") + # Plot styling constants -PLOT_WIDTH = 8 # Inches -PLOT_HEIGHT = 3.5 # Inches +PLOT_WIDTH = 8 +PLOT_HEIGHT = 3.5 + + +@dataclass +class TrainingProgress: + """Mirror of trainer.TrainingProgress — kept here so the parent process + never needs to import the heavy ML modules.""" + epoch: float = 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" + elapsed_seconds: Optional[float] = None + eta_seconds: Optional[float] = None + grad_norm: Optional[float] = None + num_tokens: Optional[int] = None + eval_loss: Optional[float] = None class TrainingBackend: """ - Training orchestration backend. - Handles both text and vision models, LoRA and full finetuning. + Training orchestration backend — subprocess-based. + Launches a fresh subprocess per training job, communicates via mp.Queue. """ def __init__(self): - self.trainer = get_trainer() + # Subprocess state + self._proc: Optional[mp.Process] = None + self._event_queue: Any = None + self._stop_queue: Any = None + self._pump_thread: Optional[threading.Thread] = None + self._lock = threading.Lock() - # 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 = [] + # Progress state (updated by pump thread from subprocess events) + self._progress = TrainingProgress() + self._should_stop = False + + # Training Metrics (consumed by routes for SSE and /metrics) + self.loss_history: list = [] + self.lr_history: list = [] + self.step_history: list = [] + self.grad_norm_history: list = [] + self.grad_norm_step_history: list = [] + self.eval_loss_history: list = [] + self.eval_step_history: list = [] + self.eval_enabled: bool = False + self.current_theme: str = "light" + + # Job metadata + self.current_job_id: Optional[str] = None + self._output_dir: Optional[str] = None + + logger.info("TrainingBackend initialized (subprocess mode)") + + # ------------------------------------------------------------------ + # Public API (called by routes/training.py) + # ------------------------------------------------------------------ + + def start_training(self, **kwargs) -> bool: + """Spawn a subprocess to run the full training pipeline. + + All kwargs are serialized into a config dict and sent to the worker. + Returns True if the subprocess was started successfully. + """ + with self._lock: + if self._proc is not None and self._proc.is_alive(): + logger.warning("Training subprocess already running") + return False + + # Join prior pump thread to prevent it from consuming events + # from the new job's queue (it reads self._event_queue dynamically). + if self._pump_thread is not None and self._pump_thread.is_alive(): + self._pump_thread.join(timeout=5.0) + if self._pump_thread.is_alive(): + logger.warning("Previous pump thread did not exit within 5s") + self._pump_thread = None + + # Reset state + self._should_stop = False + self._progress = TrainingProgress(is_training=True, status_message="Initializing training...") + self.loss_history.clear() + self.lr_history.clear() + self.step_history.clear() + self.grad_norm_history.clear() + self.grad_norm_step_history.clear() + self.eval_loss_history.clear() + self.eval_step_history.clear() self.eval_enabled = False - self.current_theme = "light" + self._output_dir = None - self.trainer.add_progress_callback(self._on_progress_update) + # Resolve project root (studio/backend/core/training/ → project root) + project_root = str(Path(__file__).resolve().parent.parent.parent.parent.parent) - logger.info("TrainingBackend initialized") + # Build config dict for the subprocess + config = { + "project_root": project_root, + "model_name": kwargs["model_name"], + "training_type": kwargs.get("training_type", "LoRA/QLoRA"), + "hf_token": kwargs.get("hf_token", ""), + "load_in_4bit": kwargs.get("load_in_4bit", True), + "max_seq_length": kwargs.get("max_seq_length", 2048), + "hf_dataset": kwargs.get("hf_dataset", ""), + "local_datasets": kwargs.get("local_datasets"), + "format_type": kwargs.get("format_type", ""), + "subset": kwargs.get("subset"), + "train_split": kwargs.get("train_split", "train"), + "eval_split": kwargs.get("eval_split"), + "eval_steps": kwargs.get("eval_steps", 0.00), + "dataset_slice_start": kwargs.get("dataset_slice_start"), + "dataset_slice_end": kwargs.get("dataset_slice_end"), + "custom_format_mapping": kwargs.get("custom_format_mapping"), + "is_dataset_multimodal": kwargs.get("is_dataset_multimodal", False), + "num_epochs": kwargs.get("num_epochs", 3), + "learning_rate": kwargs.get("learning_rate", "2e-4"), + "batch_size": kwargs.get("batch_size", 2), + "gradient_accumulation_steps": kwargs.get("gradient_accumulation_steps", 4), + "warmup_steps": kwargs.get("warmup_steps"), + "warmup_ratio": kwargs.get("warmup_ratio"), + "max_steps": kwargs.get("max_steps", 0), + "save_steps": kwargs.get("save_steps", 0), + "weight_decay": kwargs.get("weight_decay", 0.01), + "random_seed": kwargs.get("random_seed", 3407), + "packing": kwargs.get("packing", False), + "optim": kwargs.get("optim", "adamw_8bit"), + "lr_scheduler_type": kwargs.get("lr_scheduler_type", "linear"), + "use_lora": kwargs.get("use_lora", True), + "lora_r": kwargs.get("lora_r", 16), + "lora_alpha": kwargs.get("lora_alpha", 16), + "lora_dropout": kwargs.get("lora_dropout", 0.0), + "target_modules": kwargs.get("target_modules"), + "gradient_checkpointing": kwargs.get("gradient_checkpointing", "unsloth"), + "use_rslora": kwargs.get("use_rslora", False), + "use_loftq": kwargs.get("use_loftq", False), + "train_on_completions": kwargs.get("train_on_completions", False), + "finetune_vision_layers": kwargs.get("finetune_vision_layers", True), + "finetune_language_layers": kwargs.get("finetune_language_layers", True), + "finetune_attention_modules": kwargs.get("finetune_attention_modules", True), + "finetune_mlp_modules": kwargs.get("finetune_mlp_modules", True), + "enable_wandb": kwargs.get("enable_wandb", False), + "wandb_token": kwargs.get("wandb_token"), + "wandb_project": kwargs.get("wandb_project", "unsloth-training"), + "enable_tensorboard": kwargs.get("enable_tensorboard", False), + "tensorboard_dir": kwargs.get("tensorboard_dir", "runs"), + } - 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) + # Derive load_in_4bit from training_type + if config["training_type"] != "LoRA/QLoRA": + config["load_in_4bit"] = False - 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, + # Spawn subprocess + from .worker import run_training_process - # Dataset parameters - hf_dataset: str, - local_datasets: list, - format_type: str, # CHANGED: was data_template + self._event_queue = _CTX.Queue() + self._stop_queue = _CTX.Queue() - # 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, + self._proc = _CTX.Process( + target=run_training_process, + kwargs={ + "event_queue": self._event_queue, + "stop_queue": self._stop_queue, + "config": config, + }, + daemon=True, + ) + self._proc.start() + logger.info("Training subprocess started (pid=%s)", self._proc.pid) - # 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, + # Start event pump thread + self._pump_thread = threading.Thread(target=self._pump_loop, daemon=True) + self._pump_thread.start() - # 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.00, - is_dataset_multimodal: bool = False, - dataset_slice_start: int = None, - dataset_slice_end: int = None) -> 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, - eval_steps=eval_steps, - dataset_slice_start=dataset_slice_start, - dataset_slice_end=dataset_slice_end, - ) - - # 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 - - # 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 + return True def stop_training(self, save: bool = True) -> bool: - """ - Stop ongoing training. + """Send stop signal to the training subprocess.""" + self._should_stop = True + with self._lock: + if self._stop_queue is not None: + try: + self._stop_queue.put({"type": "stop", "save": save}) + except (OSError, ValueError): + pass + # Update progress immediately for responsive UI + self._progress.status_message = ( + "Stopping training and saving checkpoint..." + if save else "Cancelling training..." + ) + return True - Args: - save: If True, save the model at the current checkpoint. + def is_training_active(self) -> bool: + """Check if training is currently active.""" + with self._lock: + # Subprocess alive = active + if self._proc is not None and self._proc.is_alive(): + return True + + # Stop was requested and process exited → inactive + if self._should_stop: + return False + + # Check progress state + p = self._progress + if p.is_training: + return True + if p.is_completed or p.error: + return False + + # Check status message for activity indicators + status_lower = (p.status_message or "").lower() + if any(k in status_lower for k in ["cancelled", "canceled", "stopped", "completed", "ready to train"]): + return False + if any(k in status_lower for k in ["loading", "preparing", "training", "configuring", "tokenizing", "starting", "importing"]): + return True - 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. + """Get current training status and loss plot.""" + with self._lock: + progress = self._progress - Args: - theme: "light" or "dark" for plot styling + if not (progress.is_training or progress.is_completed or progress.error): + return (None, progress) - Returns: - Tuple of (plot, progress) - """ + plot = self._create_loss_plot(progress, theme) + return (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 - """ + def refresh_plot_for_theme(self, theme: str) -> Optional[plt.Figure]: + """Refresh plot with new theme.""" 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() + with self._lock: + progress = self._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 - """ + # ------------------------------------------------------------------ + # Compatibility shims — routes/training.py accesses these + # ------------------------------------------------------------------ + + class _TrainerShim: + """Minimal shim so routes that access backend.trainer.* still work.""" + def __init__(self, backend: "TrainingBackend"): + self._backend = backend + self.should_stop = False + + @property + def training_progress(self): + return self._backend._progress + + @training_progress.setter + def training_progress(self, value): + self._backend._progress = value + + def get_training_progress(self): + return self._backend._progress + + def _update_progress(self, **kwargs): + with self._backend._lock: + for key, value in kwargs.items(): + if hasattr(self._backend._progress, key): + setattr(self._backend._progress, key, value) + + @property + def trainer(self): + """Compatibility shim for routes that access backend.trainer.*""" + return self._TrainerShim(self) + + # ------------------------------------------------------------------ + # Event pump (background thread) + # ------------------------------------------------------------------ + + def _pump_loop(self) -> None: + """Background thread: consume events from subprocess → update state.""" + while True: + if self._proc is None or self._event_queue is None: + return + + # Try to read an event + event = self._read_queue(self._event_queue, timeout_sec=0.25) + if event is not None: + self._handle_event(event) + continue + + # No event — check if process is still alive + if self._proc.is_alive(): + continue + + # Process exited — drain remaining events + for e in self._drain_queue(self._event_queue): + self._handle_event(e) + + # Mark as done if no explicit complete/error was received + with self._lock: + if self._progress.is_training: + if self._should_stop: + self._progress.is_training = False + self._progress.status_message = "Training stopped." + else: + self._progress.is_training = False + self._progress.error = self._progress.error or "Training process exited unexpectedly" + return + + def _handle_event(self, event: dict) -> None: + """Apply a subprocess event to local state.""" + etype = event.get("type") + + with self._lock: + if etype == "progress": + self._progress.step = event.get("step", self._progress.step) + self._progress.epoch = event.get("epoch", self._progress.epoch) + self._progress.loss = event.get("loss", self._progress.loss) + self._progress.learning_rate = event.get("learning_rate", self._progress.learning_rate) + self._progress.total_steps = event.get("total_steps", self._progress.total_steps) + self._progress.elapsed_seconds = event.get("elapsed_seconds") + self._progress.eta_seconds = event.get("eta_seconds") + self._progress.grad_norm = event.get("grad_norm") + self._progress.num_tokens = event.get("num_tokens") + self._progress.eval_loss = event.get("eval_loss") + self._progress.is_training = True + status = event.get("status_message", "") + if status: + self._progress.status_message = status + + # Update metric histories + step = event.get("step", 0) + loss = event.get("loss", 0.0) + lr = event.get("learning_rate", 0.0) + if step >= 0 and loss > 0: + self.loss_history.append(loss) + self.lr_history.append(lr) + self.step_history.append(step) + + grad_norm = event.get("grad_norm") + if grad_norm is not None: + try: + gn = float(grad_norm) + except (TypeError, ValueError): + gn = None + if gn is not None and math.isfinite(gn): + self.grad_norm_history.append(gn) + self.grad_norm_step_history.append(step) + + eval_loss = event.get("eval_loss") + if eval_loss is not None: + self.eval_loss_history.append(eval_loss) + self.eval_step_history.append(step) + self.eval_enabled = True + + elif etype == "status": + self._progress.status_message = event.get("message", "") + self._progress.is_training = True + + elif etype == "complete": + self._progress.is_training = False + self._progress.is_completed = True + self._output_dir = event.get("output_dir") + msg = event.get("status_message", "Training completed") + self._progress.status_message = msg + + elif etype == "error": + self._progress.is_training = False + self._progress.error = event.get("error", "Unknown error") + logger.error("Training error: %s", event.get("error")) + stack = event.get("stack", "") + if stack: + logger.error("Stack trace:\n%s", stack) + + @staticmethod + def _read_queue(q: Any, timeout_sec: float) -> Optional[dict]: try: - training_thread = getattr(self.trainer, "training_thread", None) - if training_thread and training_thread.is_alive(): - return True + return q.get(timeout=timeout_sec) + except queue.Empty: + return None + except (EOFError, OSError, ValueError): + return None - # 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 + @staticmethod + def _drain_queue(q: Any) -> list: + events = [] + while True: + try: + events.append(q.get_nowait()) + except queue.Empty: + return events + except (EOFError, OSError, ValueError): + return events - 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 + # ------------------------------------------------------------------ + # Plot generation (unchanged from original) + # ------------------------------------------------------------------ def _create_loss_plot(self, progress: TrainingProgress, theme: str = "light") -> plt.Figure: - """ - Create training loss plot with theme-aware styling. + """Create training loss plot with theme-aware styling.""" + plt.close('all') - Args: - progress: Current training progress - theme: "light" or "dark" + 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" + } - Returns: - Matplotlib figure - """ - plt.close('all') + style = LIGHT_STYLE if theme == "light" else DARK_STYLE - # 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" - } + fig, ax = plt.subplots(figsize=(PLOT_WIDTH, PLOT_HEIGHT)) + fig.patch.set_facecolor(style["facecolor"]) + ax.set_facecolor(style["facecolor"]) - style = LIGHT_STYLE if theme == "light" else DARK_STYLE + if self.loss_history: + steps = self.step_history + losses = self.loss_history + scatter_color = "#60a5fa" + ax.scatter(steps, losses, s=16, alpha=0.6, color=scatter_color, + linewidths=0, label="Training Loss (raw)") - fig, ax = plt.subplots(figsize=(PLOT_WIDTH, PLOT_HEIGHT)) - fig.patch.set_facecolor(style["facecolor"]) - ax.set_facecolor(style["facecolor"]) + MA_WINDOW = 20 + window = min(MA_WINDOW, len(losses)) - 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)", - ) + if window >= 2: + cumsum = [0.0] + for v in losses: + cumsum.append(cumsum[-1] + float(v)) - # Moving average line overlay (trailing window) - MA_WINDOW = 20 # adjust smoothing aggressiveness - window = min(MA_WINDOW, len(losses)) + 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) - if window >= 2: - cumsum = [0.0] - for v in losses: - cumsum.append(cumsum[-1] + float(v)) + ax.plot(steps, ma, color=style["line"], linewidth=2.5, alpha=0.95, + label=f"Moving Avg ({ma[-1]:.4f})") - 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) + leg = ax.legend(frameon=False, fontsize=9) + for t in leg.get_texts(): + t.set_color(style["text"]) - ax.plot( - steps, - ma, - color=style["line"], - linewidth=2.5, - alpha=0.95, - label=f"Moving Avg ({ma[-1]:.4f})", - ) + ax.set_xlabel('Steps', fontsize=10, color=style["text"]) + ax.set_ylabel('Loss', fontsize=10, color=style["text"]) - 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"]) + 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: - 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) + title = "Training Loss" - fig.tight_layout() - return fig + ax.set_title(title, fontsize=11, fontweight='bold', pad=10, color=style["text"]) + 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 model to inference backend. + + With subprocess-based training, the model lives in the subprocess + and is freed when it exits. Inference must load from the saved + checkpoint on disk. This is a no-op placeholder. """ - 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 + logger.info( + "_transfer_to_inference_backend: subprocess training — " + "model must be loaded from disk (output_dir=%s)", self._output_dir + ) + return False # ========== GLOBAL INSTANCE ========== _training_backend = None + def get_training_backend() -> TrainingBackend: """Get global training backend instance""" global _training_backend diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py new file mode 100644 index 0000000000..e4b833cb9a --- /dev/null +++ b/studio/backend/core/training/worker.py @@ -0,0 +1,358 @@ +""" +Training subprocess entry point. + +Each training job runs in a fresh subprocess (mp.get_context("spawn")). +This gives us a clean Python interpreter with no stale module state — +solving the transformers version-switching problem completely. + +Pattern follows core/data_recipe/jobs/worker.py. +""" +from __future__ import annotations + +import logging +import os +import sys +import time +import traceback +from pathlib import Path +from typing import Any + +logger = logging.getLogger(__name__) + + +def _activate_transformers_version(model_name: str, project_root: str) -> None: + """Activate the correct transformers version BEFORE any ML imports. + + If the model needs transformers 5.x, prepend the pre-installed .venv_t5/ + directory to sys.path. Otherwise do nothing (default 4.57.x in .venv/). + """ + # Ensure backend is on path for utils imports + backend_path = os.path.join(project_root, "studio", "backend") + if backend_path not in sys.path: + sys.path.insert(0, backend_path) + + from utils.transformers_version import needs_transformers_5, _resolve_base_model + + resolved = _resolve_base_model(model_name) + if needs_transformers_5(resolved): + venv_t5 = os.path.join(project_root, ".venv_t5") + if os.path.isdir(venv_t5): + sys.path.insert(0, venv_t5) + logger.info("Activated transformers 5.x from %s", venv_t5) + else: + # Fallback: pip install at runtime (slower, ~10-15s) + logger.warning(".venv_t5 not found at %s — installing at runtime", venv_t5) + import subprocess as sp + os.makedirs(venv_t5, exist_ok=True) + r1 = sp.run( + [sys.executable, "-m", "pip", "install", "--target", venv_t5, + "--no-deps", "transformers==5.2.0"], + stdout=sp.PIPE, stderr=sp.STDOUT, + ) + r2 = sp.run( + [sys.executable, "-m", "pip", "install", "--target", venv_t5, + "--no-deps", "huggingface_hub==1.3.0"], + stdout=sp.PIPE, stderr=sp.STDOUT, + ) + if r1.returncode != 0 or r2.returncode != 0: + raise RuntimeError( + f"Failed to install transformers 5.x into {venv_t5}. " + f"pip returncode: transformers={r1.returncode}, huggingface_hub={r2.returncode}" + ) + sys.path.insert(0, venv_t5) + # Propagate to child subprocesses (e.g. GGUF converter) + _pp = os.environ.get("PYTHONPATH", "") + os.environ["PYTHONPATH"] = venv_t5 + (os.pathsep + _pp if _pp else "") + else: + logger.info("Using default transformers (4.57.x) for %s", model_name) + + +def run_training_process( + *, + event_queue: Any, + stop_queue: Any, + config: dict, +) -> None: + """Subprocess entrypoint. Fresh Python — no stale module state. + + Args: + event_queue: mp.Queue for sending progress/status/error events to parent. + stop_queue: mp.Queue for receiving stop commands from parent. + config: Training configuration dict with all parameters. + """ + os.environ["TOKENIZERS_PARALLELISM"] = "false" + + project_root = config["project_root"] + model_name = config["model_name"] + + # ── 1. Activate correct transformers version BEFORE any ML imports ── + try: + _activate_transformers_version(model_name, project_root) + except Exception as exc: + event_queue.put({ + "type": "error", + "error": f"Failed to activate transformers version: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + return + + # ── 1b. On Windows, check Triton availability (must be before import torch) ── + if sys.platform == "win32": + try: + import triton # noqa: F401 + logger.info("Triton available — torch.compile enabled") + except ImportError: + os.environ["TORCHDYNAMO_DISABLE"] = "1" + logger.warning( + "Triton not found on Windows — torch.compile disabled. " + 'Install for better performance: pip install "triton-windows<3.7"' + ) + + # ── 2. Now import ML libraries (fresh in this clean process) ── + try: + _send_status(event_queue, "Importing ML libraries...") + + backend_path = os.path.join(project_root, "studio", "backend") + if backend_path not in sys.path: + sys.path.insert(0, backend_path) + + from core.training.trainer import UnslothTrainer, TrainingProgress + + import transformers + logger.info("Subprocess loaded transformers %s", transformers.__version__) + except Exception as exc: + event_queue.put({ + "type": "error", + "error": f"Failed to import ML libraries: {exc}", + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + return + + # ── 3. Create a fresh trainer instance ── + trainer = UnslothTrainer() + + # Wire up progress callback → event_queue + def _on_progress(progress: TrainingProgress): + if progress.step >= 0 and progress.loss > 0: + event_queue.put({ + "type": "progress", + "step": progress.step, + "epoch": progress.epoch, + "loss": progress.loss, + "learning_rate": progress.learning_rate, + "total_steps": progress.total_steps, + "elapsed_seconds": progress.elapsed_seconds, + "eta_seconds": progress.eta_seconds, + "grad_norm": progress.grad_norm, + "num_tokens": progress.num_tokens, + "eval_loss": progress.eval_loss, + "status_message": progress.status_message, + "ts": time.time(), + }) + if progress.status_message: + _send_status(event_queue, progress.status_message) + + trainer.add_progress_callback(_on_progress) + + # Wire up stop_queue polling to trainer.should_stop + import threading + import queue as _queue + + def _poll_stop(): + while True: + try: + msg = stop_queue.get(timeout=1.0) + if msg and msg.get("type") == "stop": + save = msg.get("save", True) + trainer.should_stop = True + trainer.save_on_stop = save + logger.info("Stop signal received (save=%s)", save) + return + except _queue.Empty: + continue + except (EOFError, OSError): + return + + stop_thread = threading.Thread(target=_poll_stop, daemon=True) + stop_thread.start() + + # ── 4. Execute the training pipeline ── + try: + hf_token = config.get("hf_token", "") + hf_token = hf_token if hf_token and hf_token.strip() else None + + # Load model + _send_status(event_queue, "Loading model...") + success = trainer.load_model( + model_name=model_name, + max_seq_length=config["max_seq_length"], + load_in_4bit=config["load_in_4bit"], + hf_token=hf_token, + is_dataset_multimodal=config.get("is_dataset_multimodal", False), + ) + if not success or trainer.should_stop: + if trainer.should_stop: + event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()}) + else: + event_queue.put({ + "type": "error", + "error": trainer.training_progress.error or "Failed to load model", + "stack": "", "ts": time.time(), + }) + return + + # Prepare model (LoRA or full finetuning) + training_type = config.get("training_type", "LoRA/QLoRA") + use_lora = (training_type == "LoRA/QLoRA") + if use_lora: + _send_status(event_queue, "Configuring LoRA adapters...") + success = trainer.prepare_model_for_training( + use_lora=True, + finetune_vision_layers=config.get("finetune_vision_layers", True), + finetune_language_layers=config.get("finetune_language_layers", True), + finetune_attention_modules=config.get("finetune_attention_modules", True), + finetune_mlp_modules=config.get("finetune_mlp_modules", True), + target_modules=config.get("target_modules"), + lora_r=config.get("lora_r", 16), + lora_alpha=config.get("lora_alpha", 16), + lora_dropout=config.get("lora_dropout", 0.0), + use_gradient_checkpointing=config.get("gradient_checkpointing", "unsloth"), + use_rslora=config.get("use_rslora", False), + use_loftq=config.get("use_loftq", False), + ) + else: + _send_status(event_queue, "Preparing model for full finetuning...") + success = trainer.prepare_model_for_training(use_lora=False) + + if not success or trainer.should_stop: + if trainer.should_stop: + event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()}) + else: + event_queue.put({ + "type": "error", + "error": trainer.training_progress.error or "Failed to prepare model", + "stack": "", "ts": time.time(), + }) + return + + # Load dataset + _send_status(event_queue, "Loading and formatting dataset...") + hf_dataset = config.get("hf_dataset", "") + dataset_result = trainer.load_and_format_dataset( + dataset_source=hf_dataset if hf_dataset and hf_dataset.strip() else None, + format_type=config.get("format_type", ""), + local_datasets=config.get("local_datasets") or None, + custom_format_mapping=config.get("custom_format_mapping"), + subset=config.get("subset"), + train_split=config.get("train_split", "train"), + eval_split=config.get("eval_split"), + eval_steps=config.get("eval_steps", 0.00), + dataset_slice_start=config.get("dataset_slice_start"), + dataset_slice_end=config.get("dataset_slice_end"), + ) + + if isinstance(dataset_result, tuple): + dataset, eval_dataset = dataset_result + else: + dataset = dataset_result + eval_dataset = None + + # Disable eval if eval_steps <= 0 + eval_steps = config.get("eval_steps", 0.00) + if eval_steps is not None and float(eval_steps) <= 0: + eval_dataset = None + + if dataset is None or trainer.should_stop: + if trainer.should_stop: + event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()}) + else: + event_queue.put({ + "type": "error", + "error": trainer.training_progress.error or "Failed to load dataset", + "stack": "", "ts": time.time(), + }) + return + + # Convert learning rate + try: + lr_value = float(config.get("learning_rate", "2e-4")) + except ValueError: + event_queue.put({ + "type": "error", + "error": f"Invalid learning rate: {config.get('learning_rate')}", + "stack": "", "ts": time.time(), + }) + return + + # Generate output dir + output_dir = config.get("output_dir") + if not output_dir: + output_dir = f"./outputs/{model_name.replace('/', '_')}_{int(time.time())}" + + # Start training (directly — no inner thread, we ARE the subprocess) + _send_status(event_queue, "Starting training...") + max_steps = config.get("max_steps", 0) + save_steps = config.get("save_steps", 0) + + trainer._train_worker( + dataset, + output_dir=output_dir, + num_epochs=config.get("num_epochs", 3), + learning_rate=lr_value, + batch_size=config.get("batch_size", 2), + gradient_accumulation_steps=config.get("gradient_accumulation_steps", 4), + warmup_steps=config.get("warmup_steps"), + warmup_ratio=config.get("warmup_ratio"), + max_steps=max_steps if max_steps and max_steps > 0 else 0, + save_steps=save_steps if save_steps and save_steps > 0 else 0, + weight_decay=config.get("weight_decay", 0.01), + random_seed=config.get("random_seed", 3407), + packing=config.get("packing", False), + train_on_completions=config.get("train_on_completions", False), + enable_wandb=config.get("enable_wandb", False), + wandb_project=config.get("wandb_project", "unsloth-training"), + wandb_token=config.get("wandb_token"), + enable_tensorboard=config.get("enable_tensorboard", False), + tensorboard_dir=config.get("tensorboard_dir", "runs"), + eval_dataset=eval_dataset, + eval_steps=eval_steps, + max_seq_length=config.get("max_seq_length", 2048), + optim=config.get("optim", "adamw_8bit"), + lr_scheduler_type=config.get("lr_scheduler_type", "linear"), + ) + + # Check final state + progress = trainer.get_training_progress() + if progress.error: + event_queue.put({ + "type": "error", + "error": progress.error, + "stack": "", + "ts": time.time(), + }) + else: + event_queue.put({ + "type": "complete", + "output_dir": output_dir, + "status_message": progress.status_message or "Training completed", + "ts": time.time(), + }) + + except Exception as exc: + event_queue.put({ + "type": "error", + "error": str(exc), + "stack": traceback.format_exc(limit=20), + "ts": time.time(), + }) + + +def _send_status(event_queue: Any, message: str) -> None: + """Send a status update to the parent process.""" + event_queue.put({ + "type": "status", + "message": message, + "ts": time.time(), + }) diff --git a/studio/backend/main.py b/studio/backend/main.py index 331962fb67..696d8a4d10 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -3,6 +3,7 @@ Main FastAPI application for Unsloth UI Backend """ import os import secrets +import shutil from contextlib import asynccontextmanager from fastapi import FastAPI @@ -35,6 +36,12 @@ async def lifespan(app: FastAPI): # Clean up any stale compiled cache from previous runs clear_unsloth_compiled_cache() + # Remove stale .venv_overlay from previous versions — no longer used. + # Version switching now uses .venv_t5/ (pre-installed by setup.sh). + overlay_dir = Path(__file__).resolve().parent.parent.parent / ".venv_overlay" + if overlay_dir.is_dir(): + shutil.rmtree(overlay_dir, ignore_errors=True) + # Detect hardware first — sets DEVICE global used everywhere detect_hardware() @@ -89,6 +96,11 @@ app.include_router(auth_router, prefix="/api/auth", tags=["auth"]) app.include_router(training_router, prefix="/api/train", tags=["training"]) app.include_router(models_router, prefix="/api/models", tags=["models"]) app.include_router(inference_router, prefix="/api/inference", tags=["inference"]) + +# OpenAI-compatible endpoints: mount the same inference router at /v1 +# so external tools (Open WebUI, SillyTavern, etc.) can use the +# standard /v1/chat/completions path. +app.include_router(inference_router, prefix="/v1", tags=["openai-compat"]) app.include_router(datasets_router, prefix="/api/datasets", tags=["datasets"]) app.include_router(data_recipe_router, prefix="/api/data-recipe", tags=["data-recipe"]) app.include_router(export_router, prefix="/api/export", tags=["export"]) diff --git a/studio/backend/routes/export.py b/studio/backend/routes/export.py index 6616c9fbd8..7f18f50eaa 100644 --- a/studio/backend/routes/export.py +++ b/studio/backend/routes/export.py @@ -61,6 +61,42 @@ async def load_checkpoint( Wraps ExportBackend.load_checkpoint. """ try: + # Version switching is handled automatically by the subprocess-based + # export backend — no need for ensure_transformers_version() here. + + # Free GPU memory: shut down any running inference/training subprocesses + # before loading the export checkpoint (they'd compete for VRAM). + try: + from core.inference import get_inference_backend + inf = get_inference_backend() + if inf.active_model_name: + logger.info( + "Unloading inference model '%s' to free GPU memory for export", + inf.active_model_name, + ) + inf._shutdown_subprocess() + inf.active_model_name = None + inf.models.clear() + except Exception as e: + logger.warning("Could not unload inference model: %s", e) + + try: + from core.training import get_training_backend + trn = get_training_backend() + if trn.is_training_active(): + logger.info("Stopping active training to free GPU memory for export") + trn.stop_training() + # Wait for training subprocess to actually exit before proceeding, + # otherwise it may still hold GPU memory when export tries to load. + for _ in range(60): # up to 30s + if not trn.is_training_active(): + break + import time; time.sleep(0.5) + else: + logger.warning("Training subprocess did not exit within 30s, proceeding anyway") + except Exception as e: + logger.warning("Could not stop training: %s", e) + backend = get_export_backend() success, message = backend.load_checkpoint( checkpoint_path=request.checkpoint_path, diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 512eb05526..aa4b4ed0cd 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -86,6 +86,9 @@ async def load_model( GGUF models are loaded via llama-server (llama.cpp) instead of Unsloth. """ try: + # Version switching is handled automatically by the subprocess-based + # inference backend — no need for ensure_transformers_version() here. + # Create config using clean factory method # is_lora is auto-detected from adapter_config.json on disk/HF config = ModelConfig.from_identifier( @@ -160,10 +163,63 @@ async def load_model( logger.info("Unloading GGUF model before loading Unsloth model") llama_backend.unload_model() + # Shut down any export subprocess to free VRAM + try: + from core.export import get_export_backend + exp_backend = get_export_backend() + if exp_backend.current_checkpoint: + logger.info("Shutting down export subprocess to free GPU memory for inference") + exp_backend._shutdown_subprocess() + exp_backend.current_checkpoint = None + exp_backend.is_vision = False + exp_backend.is_peft = False + except Exception as e: + logger.warning("Could not shut down export subprocess: %s", e) + + # Auto-detect quantization for LoRA adapters from adapter_config.json + # The training pipeline patches this file with "unsloth_training_method" + # which is 'qlora' or 'lora'. Only LoRA (16-bit) needs load_in_4bit=False. + load_in_4bit = request.load_in_4bit + if config.is_lora and config.path: + import json + from pathlib import Path + adapter_cfg_path = Path(config.path) / "adapter_config.json" + if adapter_cfg_path.exists(): + try: + with open(adapter_cfg_path) as f: + adapter_cfg = json.load(f) + training_method = adapter_cfg.get("unsloth_training_method") + if training_method == "lora" and load_in_4bit: + logger.info( + f"adapter_config.json says unsloth_training_method='lora' — " + f"setting load_in_4bit=False to match 16-bit training" + ) + load_in_4bit = False + elif training_method == "qlora" and not load_in_4bit: + logger.info( + f"adapter_config.json says unsloth_training_method='qlora' — " + f"setting load_in_4bit=True to match QLoRA training" + ) + load_in_4bit = True + elif training_method: + logger.info(f"Training method: {training_method}, load_in_4bit={load_in_4bit}") + else: + # No unsloth_training_method — fallback to base model name + if config.base_model and "-bnb-4bit" not in config.base_model.lower() and load_in_4bit: + logger.info( + f"No unsloth_training_method in adapter_config.json. " + f"Base model '{config.base_model}' has no -bnb-4bit suffix — " + f"setting load_in_4bit=False" + ) + load_in_4bit = False + except Exception as e: + logger.warning(f"Could not read adapter_config.json: {e}") + + # Load the model success = backend.load_model( config=config, max_seq_length=request.max_seq_length, - load_in_4bit=request.load_in_4bit, + load_in_4bit=load_in_4bit, hf_token=request.hf_token, ) @@ -729,3 +785,40 @@ async def openai_chat_completions( backend.reset_generation_state() logger.error(f"Error during OpenAI completion: {e}", exc_info=True) raise HTTPException(status_code=500, detail=str(e)) + + +# ===================================================================== +# OpenAI-Compatible Models Listing (/models → /v1/models) +# ===================================================================== + +@router.get("/models") +async def openai_list_models( + current_subject: str = Depends(get_current_subject), +): + """ + OpenAI-compatible model listing endpoint. + + Returns the currently loaded model in the format expected by + OpenAI-compatible clients (``GET /v1/models``). + """ + models = [] + + # Check GGUF backend + llama_backend = get_llama_cpp_backend() + if llama_backend.is_loaded: + models.append({ + "id": llama_backend.model_identifier, + "object": "model", + "owned_by": "local", + }) + + # Check Unsloth backend + backend = get_inference_backend() + if backend.active_model_name: + models.append({ + "id": backend.active_model_name, + "object": "model", + "owned_by": "local", + }) + + return {"object": "list", "data": models} diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py index 2c7f0e846f..5f545dc133 100644 --- a/studio/backend/routes/models.py +++ b/studio/backend/routes/models.py @@ -276,8 +276,7 @@ async def get_model_config( logger.info(f"Getting model config for: {model_name}") # Load model defaults from backend config_dict = load_model_defaults(model_name) - - # Check if it's a vision model + is_vision = is_vision_model(model_name) # Check if it's a LoRA adapter diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 497daaedd3..84315291a9 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -9,7 +9,6 @@ from typing import Dict, Optional, Any import logging import asyncio from datetime import datetime -import threading # Add backend directory to path # The backend code should be in the same directory structure @@ -85,6 +84,11 @@ async def start_training( """ try: logger.info(f"Starting training job with model: {request.model_name}") + + # NOTE: No in-process ensure_transformers_version() call here. + # The subprocess (worker.py) activates the correct version in a + # fresh Python interpreter before importing any ML libraries. + backend = get_training_backend() # Generate job ID and attach to backend for later status/progress calls @@ -188,84 +192,50 @@ async def start_training( "tensorboard_dir": request.tensorboard_dir or "", } - # Set initial "preparing" state + # Free GPU memory: shut down any running inference/export subprocesses + # before training starts (they'd compete for VRAM otherwise) try: - backend.trainer._update_progress( - status_message="Initializing training...", - is_training=False, - ) - except Exception: - pass - - def run_training(): - try: + from core.inference import get_inference_backend + inf_backend = get_inference_backend() + if inf_backend.active_model_name: logger.info( - f"Starting training job {job_id} with model {request.model_name}" + "Unloading inference model '%s' to free GPU memory for training", + inf_backend.active_model_name, ) + inf_backend._shutdown_subprocess() + inf_backend.active_model_name = None + inf_backend.models.clear() + except Exception as e: + logger.warning("Could not unload inference model: %s", e) - # Update status to show we're loading model - try: - backend.trainer._update_progress(status_message="Loading model...") - except Exception as e: - logger.error(f"Error updating progress: {e}") + try: + from core.export import get_export_backend + exp_backend = get_export_backend() + if exp_backend.current_checkpoint: + logger.info("Shutting down export subprocess to free GPU memory for training") + exp_backend._shutdown_subprocess() + exp_backend.current_checkpoint = None + exp_backend.is_vision = False + exp_backend.is_peft = False + except Exception as e: + logger.warning("Could not shut down export subprocess: %s", e) - # start_training returns bool (not generator) - run_result = backend.start_training(**training_kwargs) - logger.info( - "Training job %s backend.start_training returned type=%s value=%r", - job_id, - type(run_result).__name__, - run_result, - ) - if not run_result: - progress_error = backend.trainer.training_progress.error - raise RuntimeError(progress_error or "Training failed to start") + # start_training now spawns a subprocess (non-blocking) + success = backend.start_training(**training_kwargs) - logger.info(f"Training job {job_id} started successfully") - - except Exception as e: - logger.error(f"Training error in job {job_id}: {e}", exc_info=True) - try: - backend.trainer._update_progress( - error=str(e), - is_training=False, - ) - except Exception as update_error: - logger.error(f"Failed to update progress: {update_error}") - - # Start training in a daemon thread - training_thread = threading.Thread( - target=run_training, - daemon=True, - name=f"Training-{job_id}", - ) - training_thread.start() - - # Store thread reference for status checking - backend._training_thread = training_thread - - # Give it a moment to start - import time - - time.sleep(0.5) - - # Verify training thread is alive - if not training_thread.is_alive(): - logger.warning(f"Training thread died immediately for job {job_id}") + if not success: + progress_error = backend.trainer.training_progress.error return TrainingJobResponse( job_id=job_id, status="error", - message=( - "Training thread failed to start. " - "Check server logs for details." - ), - error="Thread not alive", + message=progress_error or "Failed to start training subprocess", + error=progress_error or "subprocess_start_failed", ) return TrainingJobResponse( job_id=job_id, status="queued", - message="Training job queued and starting in background", + message="Training job queued and starting in subprocess", error=None, ) @@ -290,18 +260,10 @@ async def stop_training( """ try: backend = get_training_backend() - trainer_thread = getattr(getattr(backend, "trainer", None), "training_thread", None) - thread_alive = bool(trainer_thread and trainer_thread.is_alive()) is_active = backend.is_training_active() - logger.info( - "Stop requested: save=%s is_active=%s thread_alive=%s should_stop=%s", - body.save, - is_active, - thread_alive, - getattr(getattr(backend, "trainer", None), "should_stop", None), - ) + logger.info("Stop requested: save=%s is_active=%s", body.save, is_active) - if not is_active and not thread_alive: + if not is_active: return TrainingStopResponse( status="idle", message="No training job is currently running" @@ -332,25 +294,22 @@ async def reset_training( """ try: backend = get_training_backend() - trainer_thread = getattr(getattr(backend, "trainer", None), "training_thread", None) - thread_alive = bool(trainer_thread and trainer_thread.is_alive()) is_active = backend.is_training_active() - if is_active or thread_alive: - logger.warning( - "Rejected reset while training active: is_active=%s thread_alive=%s should_stop=%s", - is_active, - thread_alive, - getattr(getattr(backend, "trainer", None), "should_stop", None), - ) + if is_active: + logger.warning("Rejected reset while training active: is_active=%s", is_active) raise HTTPException( status_code=409, detail="Training is still running. Stop training and wait for it to finish before resetting.", ) logger.info("Reset training state: clearing runtime + metric history") - backend.trainer.should_stop = False - backend.trainer.training_progress = backend.trainer.training_progress.__class__() + backend._should_stop = False # Clear stop flag so status returns to idle + backend.trainer._update_progress( + is_training=False, is_completed=False, error=None, + status_message="Ready to train", step=0, loss=0.0, epoch=0, + total_steps=0, + ) backend.loss_history = [] backend.lr_history = [] backend.step_history = [] @@ -376,18 +335,11 @@ async def get_training_status( """ try: backend = get_training_backend() - job_id: str = getattr(backend, "current_job_id", "") + job_id: str = getattr(backend, "current_job_id", "") or "" # Check if training is active is_active = backend.is_training_active() - # Check if there's a training thread running (preparation phase) - has_thread = ( - hasattr(backend, "_training_thread") - and backend._training_thread - and backend._training_thread.is_alive() - ) - # Get progress info from trainer try: progress = backend.trainer.get_training_progress() @@ -400,14 +352,14 @@ async def get_training_status( error_message = getattr(progress, "error", None) if progress else None # Check if training was stopped by user - trainer_stopped = getattr(backend.trainer, "should_stop", False) + trainer_stopped = getattr(backend, "_should_stop", False) # Derive high-level phase if error_message: phase = "error" elif is_active: msg_lower = status_message.lower() - if "loading" in msg_lower: + if "loading" in msg_lower or "importing" in msg_lower: phase = "loading_model" elif any( k in msg_lower for k in ["preparing", "initializing", "configuring"] @@ -419,8 +371,6 @@ async def get_training_status( phase = "stopped" elif progress and getattr(progress, "is_completed", False): phase = "completed" - elif has_thread: - phase = "loading_model" else: phase = "idle" @@ -534,7 +484,7 @@ async def stream_training_progress( async def event_generator(): backend = get_training_backend() - job_id: str = getattr(backend, "current_job_id", "") + job_id: str = getattr(backend, "current_job_id", "") or "" # ── Helpers ────────────────────────────────────────────── def build_progress( diff --git a/studio/backend/utils/models/model_config.py b/studio/backend/utils/models/model_config.py index 5404a198aa..e7b92651be 100644 --- a/studio/backend/utils/models/model_config.py +++ b/studio/backend/utils/models/model_config.py @@ -7,6 +7,9 @@ from typing import Optional, Dict, Any from utils.paths import normalize_path, is_local_path, is_model_cached from utils.utils import without_hf_auth import logging +import os +import subprocess +import sys from pathlib import Path from typing import List, Tuple import json @@ -364,7 +367,117 @@ def load_model_config(model_name: str, use_auth: bool = False, token: Optional[s model_name, trust_remote_code=True ) -pass + + +# VLM architecture suffixes and known VLM model_type values. +_VLM_ARCH_SUFFIXES = ("ForConditionalGeneration", "ForVisionText2Text") +_VLM_MODEL_TYPES = { + 'phi3_v', 'llava', 'llava_next', 'llava_onevision', + 'internvl_chat', 'cogvlm2', 'minicpmv', +} + +# Pre-computed project root and .venv_t5 path for subprocess version switching. +_PROJECT_ROOT = str(Path(__file__).resolve().parent.parent.parent.parent.parent) +_VENV_T5_DIR = os.path.join(_PROJECT_ROOT, ".venv_t5") +_BACKEND_DIR = os.path.join(_PROJECT_ROOT, "studio", "backend") + +# Inline script executed in a subprocess with transformers 5.x activated. +# Receives model_name and token via argv, prints JSON result to stdout. +_VISION_CHECK_SCRIPT = r''' +import sys, os, json +os.environ["TOKENIZERS_PARALLELISM"] = "false" + +# Activate transformers 5.x +venv_t5 = sys.argv[1] +backend_dir = sys.argv[2] +model_name = sys.argv[3] +token = sys.argv[4] if len(sys.argv) > 4 and sys.argv[4] != "" else None + +sys.path.insert(0, venv_t5) +if backend_dir not in sys.path: + sys.path.insert(0, backend_dir) + +try: + from transformers import AutoConfig + kwargs = {"trust_remote_code": True} + if token: + kwargs["token"] = token + config = AutoConfig.from_pretrained(model_name, **kwargs) + + is_vlm = False + if hasattr(config, "architectures"): + is_vlm = any( + x.endswith(("ForConditionalGeneration", "ForVisionText2Text")) + for x in config.architectures + ) + if not is_vlm and hasattr(config, "vision_config"): + is_vlm = True + if not is_vlm and hasattr(config, "img_processor"): + is_vlm = True + if not is_vlm and hasattr(config, "image_token_index"): + is_vlm = True + if not is_vlm and hasattr(config, "model_type"): + vlm_types = {"phi3_v","llava","llava_next","llava_onevision", + "internvl_chat","cogvlm2","minicpmv"} + if config.model_type in vlm_types: + is_vlm = True + + model_type = getattr(config, "model_type", "unknown") + archs = getattr(config, "architectures", []) + print(json.dumps({"is_vision": is_vlm, "model_type": model_type, + "architectures": archs})) +except Exception as exc: + print(json.dumps({"error": str(exc)})) + sys.exit(1) +''' + + +def _is_vision_model_subprocess(model_name: str, hf_token: Optional[str] = None) -> bool: + """Run is_vision_model check in a subprocess with transformers 5.x. + + Same pattern as training/inference workers: spawn a clean subprocess + with .venv_t5/ prepended to sys.path so AutoConfig recognizes newer + architectures (glm4_moe_lite, etc.). + """ + token_arg = hf_token or "" + + try: + result = subprocess.run( + [sys.executable, "-c", _VISION_CHECK_SCRIPT, + _VENV_T5_DIR, _BACKEND_DIR, model_name, token_arg], + capture_output=True, text=True, timeout=60, + ) + + if result.returncode != 0: + stderr = result.stderr.strip() + logger.warning( + "Vision check subprocess failed for '%s': %s", + model_name, stderr or result.stdout.strip(), + ) + return False + + data = json.loads(result.stdout.strip()) + if "error" in data: + logger.warning( + "Vision check subprocess error for '%s': %s", + model_name, data["error"], + ) + return False + + is_vlm = data["is_vision"] + logger.info( + "Vision check (subprocess, transformers 5.x) for '%s': " + "model_type=%s, architectures=%s, is_vision=%s", + model_name, data.get("model_type"), data.get("architectures"), is_vlm, + ) + return is_vlm + + except subprocess.TimeoutExpired: + logger.warning("Vision check subprocess timed out for '%s'", model_name) + return False + except Exception as exc: + logger.warning("Vision check subprocess failed for '%s': %s", model_name, exc) + return False def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool: @@ -372,17 +485,32 @@ def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool: Detect vision-language models (VLMs) by checking architecture in config. Works for fine-tuned models since they inherit the base architecture. + For models that require transformers 5.x (e.g. GLM-4.7-Flash), the check + runs in a subprocess with .venv_t5/ activated — same pattern as the + training and inference workers. + Args: model_name: Model identifier (HF repo or local path) hf_token: Optional HF token for accessing gated/private models """ + # Models that need transformers 5.x must be checked in a subprocess + # because AutoConfig in the main process (transformers 4.57.x) doesn't + # recognize their architectures. + from utils.transformers_version import needs_transformers_5 + if needs_transformers_5(model_name): + logger.info( + "Model '%s' needs transformers 5.x — checking vision via subprocess", + model_name, + ) + return _is_vision_model_subprocess(model_name, hf_token=hf_token) + try: config = load_model_config(model_name, use_auth=True, token=hf_token) # Check 1: Architecture class name patterns if hasattr(config, 'architectures'): is_vlm = any( - x.endswith(("ForConditionalGeneration", "ForVisionText2Text")) + x.endswith(_VLM_ARCH_SUFFIXES) for x in config.architectures ) if is_vlm: @@ -406,11 +534,7 @@ def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool: # Check 5: Known VLM model_type values that may not match above checks if hasattr(config, 'model_type'): - vlm_model_types = { - 'phi3_v', 'llava', 'llava_next', 'llava_onevision', - 'internvl_chat', 'cogvlm2', 'minicpmv', - } - if config.model_type in vlm_model_types: + if config.model_type in _VLM_MODEL_TYPES: logger.info(f"Model {model_name} detected as VLM: model_type={config.model_type}") return True @@ -419,7 +543,6 @@ def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool: except Exception as e: logger.warning(f"Could not determine if {model_name} is vision model: {e}") return False -pass def _is_mmproj(filename: str) -> bool: diff --git a/studio/backend/utils/transformers_version.py b/studio/backend/utils/transformers_version.py new file mode 100644 index 0000000000..8efbab46a2 --- /dev/null +++ b/studio/backend/utils/transformers_version.py @@ -0,0 +1,266 @@ +""" +Automatic transformers version switching. + +Some newer model architectures (Ministral-3, GLM-4.7-Flash, Qwen3-30B-A3B MoE, +tiny_qwen3_moe) require transformers>=5.2.0, while everything else needs the +default 4.57.x that ships with Unsloth. + +When loading a LoRA adapter with a custom name, we resolve the base model from +``adapter_config.json`` and check *that* against the model list. + +Strategy: + Training and inference run in subprocesses that activate the correct version + via sys.path (prepending .venv_t5/ for 5.x models). See: + - core/training/worker.py + - core/inference/worker.py + + For export (still in-process), ensure_transformers_version() does a lightweight + sys.path swap using the same .venv_t5/ directory pre-installed by setup.sh. +""" + +import importlib +import json +import logging +import os +import subprocess +import sys +from pathlib import Path + +logger = logging.getLogger(__name__) + +# Ensure our logger is visible even if root logger isn't configured for INFO. +if not logger.handlers: + _handler = logging.StreamHandler() + _handler.setLevel(logging.INFO) + _handler.setFormatter( + logging.Formatter("[%(name)s|%(levelname)s]%(message)s") + ) + logger.addHandler(_handler) + logger.setLevel(logging.INFO) + +# --------------------------------------------------------------------------- +# Detection +# --------------------------------------------------------------------------- + +# Lowercase substrings — if ANY appears anywhere in the lowered model name, +# we need transformers 5.x. +TRANSFORMERS_5_MODEL_SUBSTRINGS: tuple[str, ...] = ( + "ministral-3-", # Ministral-3-{3,8,14}B-{Instruct,Reasoning,Base}-2512 + "glm-4.7-flash", # GLM-4.7-Flash + "qwen3-30b-a3b", # Qwen3-30B-A3B-Instruct-2507 and variants + "qwen3.5", # Qwen3.5 family (35B-A3B, etc.) + "qwen3-next", # Qwen3-Next and variants + "tiny_qwen3_moe", # imdatta0/tiny_qwen3_moe_2.8B_0.7B +) + +# Versions +TRANSFORMERS_5_VERSION = "5.2.0" +TRANSFORMERS_DEFAULT_VERSION = "4.57.1" + +# Pre-installed directory for transformers 5.x — created by setup.sh / setup.ps1 +_PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent.parent # studio/backend/utils/ → project root +_VENV_T5_DIR = str(_PROJECT_ROOT / ".venv_t5") + + +def _resolve_base_model(model_name: str) -> str: + """If *model_name* points to a LoRA adapter, return its base model. + + Checks for ``adapter_config.json`` locally first. Only calls the heavier + ``get_base_model_from_lora`` for paths that are actual local directories + (avoids noisy warnings for plain HF model IDs). + + Returns the original *model_name* unchanged if it is not a LoRA adapter. + """ + # --- Fast local check --------------------------------------------------- + local_path = Path(model_name) + adapter_cfg_path = local_path / "adapter_config.json" + if adapter_cfg_path.is_file(): + try: + with open(adapter_cfg_path) as f: + cfg = json.load(f) + base = cfg.get("base_model_name_or_path") + if base: + logger.info( + "Resolved LoRA adapter '%s' → base model '%s'", + model_name, base, + ) + return base + except Exception as exc: + logger.debug("Could not read %s: %s", adapter_cfg_path, exc) + + # --- Only try the heavier fallback for local directories ---------------- + if local_path.is_dir(): + try: + from utils.models import get_base_model_from_lora + base = get_base_model_from_lora(model_name) + if base: + logger.info( + "Resolved LoRA adapter '%s' → base model '%s' " + "(via get_base_model_from_lora)", + model_name, base, + ) + return base + except Exception as exc: + logger.debug( + "get_base_model_from_lora failed for '%s': %s", + model_name, exc, + ) + + return model_name + + +def needs_transformers_5(model_name: str) -> bool: + """Return True if *model_name* belongs to an architecture that requires + ``transformers>=5.2.0``.""" + lowered = model_name.lower() + return any(sub in lowered for sub in TRANSFORMERS_5_MODEL_SUBSTRINGS) + + +# --------------------------------------------------------------------------- +# Version switching (in-process — used only by export) +# --------------------------------------------------------------------------- + +def _get_in_memory_version() -> str | None: + """Return the transformers version currently loaded in this process.""" + tf = sys.modules.get("transformers") + if tf is not None: + return getattr(tf, "__version__", None) + return None + + +# All top-level prefixes that hold references to transformers internals. +_PURGE_PREFIXES = ( + "transformers", + "huggingface_hub", + "unsloth", + "unsloth_zoo", + "peft", + "trl", + "accelerate", + "auto_gptq", + # NOTE: bitsandbytes is intentionally EXCLUDED — it registers torch custom + # operators at import time via torch.library.define(). Those registrations + # live in torch's global operator registry which survives module purge. + # Re-importing bitsandbytes after purge → duplicate registration → crash. + # Our own modules that import from transformers at module level + # (e.g. model_config.py: `from transformers import AutoConfig`) + "utils.models", + "core.training", + "core.inference", + "core.export", +) + + +def _purge_modules() -> int: + """Remove all cached modules for transformers and its dependents. + + Returns the number of modules purged. + """ + importlib.invalidate_caches() + to_remove = [ + k for k in list(sys.modules.keys()) + if any(k == p or k.startswith(p + ".") for p in _PURGE_PREFIXES) + ] + for key in to_remove: + del sys.modules[key] + return len(to_remove) + + +def _ensure_venv_t5_exists() -> bool: + """Ensure .venv_t5/ exists. Install at runtime if missing.""" + if os.path.isdir(_VENV_T5_DIR) and os.listdir(_VENV_T5_DIR): + return True + + logger.warning(".venv_t5 not found at %s — installing at runtime", _VENV_T5_DIR) + os.makedirs(_VENV_T5_DIR, exist_ok=True) + for pkg in (f"transformers=={TRANSFORMERS_5_VERSION}", "huggingface_hub==1.3.0"): + cmd = [ + sys.executable, "-m", "pip", "install", + "--target", _VENV_T5_DIR, + "--no-deps", + pkg, + ] + result = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True) + if result.returncode != 0: + logger.error("pip install failed:\n%s", result.stdout) + return False + logger.info("Installed transformers 5.x to %s", _VENV_T5_DIR) + return True + + +def _activate_5x() -> None: + """Prepend .venv_t5/ to sys.path, purge stale modules, reimport.""" + if not _ensure_venv_t5_exists(): + raise RuntimeError(f"Cannot activate transformers 5.x: .venv_t5 missing at {_VENV_T5_DIR}") + + if _VENV_T5_DIR not in sys.path: + sys.path.insert(0, _VENV_T5_DIR) + logger.info("Prepended %s to sys.path", _VENV_T5_DIR) + + count = _purge_modules() + logger.info("Purged %d cached modules", count) + + import transformers + logger.info("Loaded transformers %s", transformers.__version__) + + +def _deactivate_5x() -> None: + """Remove .venv_t5/ from sys.path, purge stale modules, reimport.""" + while _VENV_T5_DIR in sys.path: + sys.path.remove(_VENV_T5_DIR) + logger.info("Removed %s from sys.path", _VENV_T5_DIR) + + count = _purge_modules() + logger.info("Purged %d cached modules", count) + + import transformers + logger.info("Reverted to transformers %s", transformers.__version__) + + +def ensure_transformers_version(model_name: str) -> None: + """Ensure the correct ``transformers`` version is active for *model_name*. + + Uses sys.path with .venv_t5/ (pre-installed by setup.sh): + • Need 5.x → prepend .venv_t5/ to sys.path, purge modules. + • Need 4.x → remove .venv_t5/ from sys.path, purge modules. + + For LoRA adapters with custom names, the base model is resolved from + ``adapter_config.json`` before checking. + + NOTE: Training and inference use subprocess isolation instead of this + function. This is only used by the export path (routes/export.py). + """ + # Resolve LoRA adapters to their base model for accurate detection + resolved = _resolve_base_model(model_name) + want_5 = needs_transformers_5(resolved) + target_version = TRANSFORMERS_5_VERSION if want_5 else TRANSFORMERS_DEFAULT_VERSION + target_major = int(target_version.split(".")[0]) + + # Check what's actually loaded in memory + in_memory = _get_in_memory_version() + + logger.info( + "Version check for '%s' (resolved: '%s'): need=%s, in_memory=%s", + model_name, resolved, target_version, in_memory, + ) + + # --- Already correct? --------------------------------------------------- + if in_memory is not None: + in_memory_major = int(in_memory.split(".")[0]) + if in_memory_major == target_major: + logger.info( + "transformers %s already loaded — correct for '%s'", + in_memory, model_name, + ) + return + + # --- Switch version ----------------------------------------------------- + if want_5: + logger.info("Activating transformers %s via .venv_t5…", TRANSFORMERS_5_VERSION) + _activate_5x() + else: + logger.info("Reverting to default transformers %s…", TRANSFORMERS_DEFAULT_VERSION) + _deactivate_5x() + + final = _get_in_memory_version() + logger.info("✓ transformers version is now %s", final) diff --git a/studio/frontend/src/features/chat/api/chat-api.ts b/studio/frontend/src/features/chat/api/chat-api.ts index 5d5a9551ef..4276e618ed 100644 --- a/studio/frontend/src/features/chat/api/chat-api.ts +++ b/studio/frontend/src/features/chat/api/chat-api.ts @@ -99,7 +99,7 @@ export async function* streamChatCompletions( payload: OpenAIChatCompletionsRequest, signal: AbortSignal, ): AsyncGenerator { - const response = await authFetch("/api/inference/chat/completions", { + const response = await authFetch("/v1/chat/completions", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify(payload), diff --git a/studio/frontend/src/features/chat/runtime-provider.tsx b/studio/frontend/src/features/chat/runtime-provider.tsx index ce553dcdbf..46e3c697ad 100644 --- a/studio/frontend/src/features/chat/runtime-provider.tsx +++ b/studio/frontend/src/features/chat/runtime-provider.tsx @@ -202,7 +202,7 @@ async function generateTitleWithModel(payload: { return joined.length > 60 ? joined.slice(0, 60).trimEnd() : joined; } - const response = await authFetch("/api/inference/chat/completions", { + const response = await authFetch("/v1/chat/completions", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ diff --git a/studio/frontend/vite.config.ts b/studio/frontend/vite.config.ts index 64ff899436..c8c9028b2e 100644 --- a/studio/frontend/vite.config.ts +++ b/studio/frontend/vite.config.ts @@ -17,6 +17,10 @@ export default defineConfig({ target: "http://127.0.0.1:8000", changeOrigin: true, }, + "/v1": { + target: "http://127.0.0.1:8000", + changeOrigin: true, + }, "/seed/inspect": { target: "http://127.0.0.1:8004", changeOrigin: true,