168 lines
5.9 KiB
Python
168 lines
5.9 KiB
Python
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)
|