import whisper import os import sys import subprocess import torch import tracker def check_gpu_health(): """ Performs a robust check for GPU availability and prints detailed troubleshooting info if issues are detected, specific to Bazzite/VS Code environments. """ tracker.logger.info("Checking GPU health...") # 1. Check if the OS/Driver sees the GPU nvidia_smi_ok = False try: in_flatpak = os.path.exists("/.flatpak-info") cmd = ["flatpak-spawn", "--host", "nvidia-smi"] if in_flatpak else ["nvidia-smi"] subprocess.run(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True) nvidia_smi_ok = True except (subprocess.CalledProcessError, FileNotFoundError): nvidia_smi_ok = False # 2. Check if PyTorch sees the GPU torch_cuda_ok = torch.cuda.is_available() if torch_cuda_ok: tracker.logger.info(f"✅ GPU is accessible: {torch.cuda.get_device_name(0)}") tracker.logger.info(f" CUDA Version: {torch.version.cuda}") return True # --- Troubleshooting Block --- tracker.logger.warning("\n⚠️ WARNING: GPU not detected by PyTorch. Falling back to CPU.") tracker.logger.warning(" Transcription will be significantly slower.\n") tracker.logger.info("--- Diagnostic Report ---") if nvidia_smi_ok: tracker.logger.info("1. [OK] 'nvidia-smi' command works. The system driver is installed and visible.") tracker.logger.info("2. [FAIL] PyTorch cannot see the GPU.") tracker.logger.info(" -> Likely Cause: You might have installed the CPU-only version of PyTorch.") tracker.logger.info(" -> Solution: Reinstall PyTorch with CUDA support:") tracker.logger.info(" pip uninstall torch torchvision torchaudio") tracker.logger.info(" pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118") else: tracker.logger.info("1. [FAIL] 'nvidia-smi' command failed or not found.") tracker.logger.info(" -> Likely Cause: Nvidia drivers are missing, or the container/sandbox cannot access the GPU.") tracker.logger.info("\n --- Bazzite / VS Code / Container Specific Checks ---") tracker.logger.info(" a. If you are running inside a dev container (DevBox/Distrobox/Toolbox):") tracker.logger.info(" Ensure the container was created with nvidia support.") tracker.logger.info(" (Bazzite usually handles this for 'distrobox', but check your config).") tracker.logger.info(" b. If you are using VS Code Flatpak:") tracker.logger.info(" Flatpak might be restricting access. Check Flatseal permissions for VS Code.") tracker.logger.info(" c. Driver Check:") tracker.logger.info(" Run 'rpm -qa | grep nvidia' in your host terminal to verify drivers are installed.") tracker.logger.info("-------------------------\n") return False def get_vram_gb(): """Returns the total VRAM in GB of the first CUDA device, or 0 if no CUDA.""" if not torch.cuda.is_available(): return 0 try: # returns bytes total_mem = torch.cuda.get_device_properties(0).total_memory return total_mem / (1024 ** 3) except Exception: return 0 def get_optimal_model_size(): """ Determines the best Whisper model based on available VRAM. Rough estimates for VRAM usage (fp16): - large: ~10 GB - medium: ~5 GB - small: ~2 GB - base: ~1 GB - tiny: ~1 GB """ vram = get_vram_gb() if vram == 0: return "base" tracker.logger.info(f"Detected GPU with {vram:.2f} GB VRAM.") if vram >= 11: return "large" elif vram >= 6: return "medium" elif vram >= 3: return "small" else: return "base" def format_timestamp(seconds: float): """Converts seconds to SRT timestamp format (HH:MM:SS,mmm).""" whole_seconds = int(seconds) milliseconds = int((seconds - whole_seconds) * 1000) hours = whole_seconds // 3600 minutes = (whole_seconds % 3600) // 60 seconds = whole_seconds % 60 return f"{hours:02d}:{minutes:02d}:{seconds:02d},{milliseconds:03d}" def save_as_srt(result, output_path): """Saves the Whisper transcription result as an SRT file.""" with open(output_path, "w", encoding="utf-8") as f: for i, segment in enumerate(result["segments"], start=1): start = format_timestamp(segment["start"]) end = format_timestamp(segment["end"]) text = segment["text"].strip() f.write(f"{i}\n") f.write(f"{start} --> {end}\n") f.write(f"{text}\n\n") tracker.logger.info(f"SRT saved to: {output_path}") def load_whisper_model(model_size="auto"): """ Loads and returns the Whisper model. """ check_gpu_health() if model_size == "auto": model_size = get_optimal_model_size() tracker.logger.info(f"Auto-selected model: '{model_size}'") tracker.logger.info(f"Loading Whisper model ('{model_size}')...") device = "cuda" if torch.cuda.is_available() else "cpu" tracker.logger.info(f"Using device: {device}") try: model = whisper.load_model(model_size, device=device) return model except Exception as e: tracker.logger.error(f"Error loading model: {e}") sys.exit(1) def transcribe_audio(audio_path, model_size="auto", language=None, loaded_model=None): """ Transcribes an audio file using OpenAI's Whisper model. Args: audio_path (str): Path to the input audio file. model_size (str): Size of the Whisper model to use. If "auto", selects based on VRAM. language (str, optional): Language code (e.g., "en", "fr", "es"). If None, auto-detects. loaded_model (object, optional): Pre-loaded Whisper model object. Returns: dict: The full transcription result containing segments and text. """ if not os.path.exists(audio_path): raise FileNotFoundError(f"Audio file not found: {audio_path}") model = loaded_model if model is None: model = load_whisper_model(model_size) tracker.logger.info(f"Transcribing {audio_path}...") try: # Disable verbose to prevent line-by-line output result = model.transcribe(audio_path, language=language, verbose=False) tracker.logger.info("Transcription complete.") return result except Exception as e: tracker.logger.error(f"Error during transcription: {e}") sys.exit(1)