diff --git a/studio/backend/core/export/export.py b/studio/backend/core/export/export.py index 11662c9e5f..8977586ac3 100644 --- a/studio/backend/core/export/export.py +++ b/studio/backend/core/export/export.py @@ -11,6 +11,7 @@ from unsloth import FastLanguageModel, FastVisionModel from huggingface_hub import HfApi, ModelCard from transformers.modeling_utils import PushToHubMixin import torch +from utils.hardware import clear_gpu_cache from utils.models import is_vision_model, get_base_model_from_lora from core.inference import get_inference_backend @@ -69,14 +70,8 @@ class ExportBackend: self.current_tokenizer = None self.current_checkpoint = None - # Force garbage collection - import gc - gc.collect() - - # Clear CUDA cache - if torch.cuda.is_available(): - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + # Clear GPU memory cache (handles gc + backend-specific cleanup) + clear_gpu_cache() logger.info("Memory cleanup completed successfully") return True diff --git a/studio/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py index 70b6f49f32..fe75d6a976 100644 --- a/studio/backend/core/inference/inference.py +++ b/studio/backend/core/inference/inference.py @@ -11,7 +11,8 @@ import torch from typing import Optional, Generator, Tuple from utils.models import ModelConfig, get_base_model_from_lora from utils.paths import is_model_cached -from utils.utils import format_error_message, log_gpu_memory +from utils.utils import format_error_message +from utils.hardware import get_device, clear_gpu_cache, log_gpu_memory from io import StringIO import logging @@ -35,7 +36,7 @@ class InferenceBackend: "unsloth/Gemma-3-4B-it", "unsloth/Qwen2-VL-2B-Instruct-bnb-4bit", ] - self.device = "cuda" if torch.cuda.is_available() else "cpu" + self.device = get_device().value # Thread safety import threading @@ -154,12 +155,8 @@ class InferenceBackend: if self.active_model_name == model_name: self.active_model_name = None - # Use garbage collection and clear CUDA cache to release memory - import gc - import torch - gc.collect() - if torch.cuda.is_available(): - torch.cuda.empty_cache() + # Clear GPU memory cache + clear_gpu_cache() logger.info(f"Model '{model_name}' successfully unloaded.") return True @@ -562,11 +559,11 @@ class InferenceBackend: input_text, add_special_tokens=False, return_tensors="pt", - ).to("cuda") + ).to(self.device) else: # Text-only for vision model formatted_prompt = self.format_chat_prompt(messages, system_prompt) - inputs = processor.tokenizer(formatted_prompt, return_tensors="pt").to("cuda") + inputs = processor.tokenizer(formatted_prompt, return_tensors="pt").to(self.device) # Generate with streaming captured_output = StringIO() @@ -888,11 +885,8 @@ class InferenceBackend: for model_name in self.models.keys(): self._reset_model_generation_state(model_name) - import torch - if torch.cuda.is_available(): - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - logger.debug("Cleared CUDA cache and IPC resources") + clear_gpu_cache() + logger.debug("Cleared GPU cache") import gc gc.collect() diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 5fc48aa18a..af5c3ebe01 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 Gradio UI """ import torch +from utils.hardware import clear_gpu_cache torch._dynamo.config.recompile_limit = 64 from unsloth import FastLanguageModel, FastVisionModel, is_bfloat16_supported from unsloth.chat_templates import get_chat_template @@ -98,9 +99,7 @@ class UnslothTrainer: """Load model for training (supports both text and vision models)""" try: print("\nClearing GPU memory before training...") - torch.cuda.empty_cache() - import gc - gc.collect() + clear_gpu_cache() # Detect if this is a vision model first self.is_vlm = is_vision_model(model_name) @@ -804,8 +803,7 @@ class UnslothTrainer: self.tokenizer = None # Clear GPU memory - if torch.cuda.is_available(): - torch.cuda.empty_cache() + clear_gpu_cache() def _ensure_deepseek_ocr_installed():