import whisper import os import sys import subprocess import torch 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. """ print("Checking GPU health...") # 1. Check if the OS/Driver sees the GPU nvidia_smi_ok = False try: subprocess.run(["nvidia-smi"], 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: print(f"✅ GPU is accessible: {torch.cuda.get_device_name(0)}") print(f" CUDA Version: {torch.version.cuda}") return True # --- Troubleshooting Block --- print("\n⚠️ WARNING: GPU not detected by PyTorch. Falling back to CPU.") print(" Transcription will be significantly slower.\n") print("--- Diagnostic Report ---") if nvidia_smi_ok: print("1. [OK] 'nvidia-smi' command works. The system driver is installed and visible.") print("2. [FAIL] PyTorch cannot see the GPU.") print(" -> Likely Cause: You might have installed the CPU-only version of PyTorch.") print(" -> Solution: Reinstall PyTorch with CUDA support:") print(" pip uninstall torch torchvision torchaudio") print(" pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118") else: print("1. [FAIL] 'nvidia-smi' command failed or not found.") print(" -> Likely Cause: Nvidia drivers are missing, or the container/sandbox cannot access the GPU.") print("\n --- Bazzite / VS Code / Container Specific Checks ---") print(" a. If you are running inside a dev container (DevBox/Distrobox/Toolbox):") print(" Ensure the container was created with nvidia support.") print(" (Bazzite usually handles this for 'distrobox', but check your config).") print(" b. If you are using VS Code Flatpak:") print(" Flatpak might be restricting access. Check Flatseal permissions for VS Code.") print(" c. Driver Check:") print(" Run 'rpm -qa | grep nvidia' in your host terminal to verify drivers are installed.") print("-------------------------\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" print(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") print(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() print(f"Auto-selected model: '{model_size}'") print(f"Loading Whisper model ('{model_size}')...") device = "cuda" if torch.cuda.is_available() else "cpu" print(f"Using device: {device}") try: model = whisper.load_model(model_size, device=device) return model except Exception as e: print(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) print(f"Transcribing {audio_path}...") try: # Enable verbose=True to show progress in terminal result = model.transcribe(audio_path, language=language, verbose=True) print("Transcription complete.") return result except Exception as e: print(f"Error during transcription: {e}") sys.exit(1)