diff --git a/.env.example b/.env.example index 9d73362b..022f08a7 100644 --- a/.env.example +++ b/.env.example @@ -38,7 +38,7 @@ WHISPER_MODEL=base HF_TOKEN="" -# WHISPER_DEVICE: "cpu" | "cuda" | "auto" (auto = try GPU, fall back to CPU) +# WHISPER_DEVICE: "cpu" | "cuda" WHISPER_DEVICE=cpu diff --git a/README.md b/README.md index 27bc732a..4bc189d7 100644 --- a/README.md +++ b/README.md @@ -248,8 +248,8 @@ uv sync --extra voice | Variable | Description | Default | |----------|-------------|---------| | `VOICE_NOTE_ENABLED` | Enable voice note handling | `true` | -| `WHISPER_MODEL` | Model size: `tiny`, `base`, `small`, `medium`, `large-v2` | `base` | -| `WHISPER_DEVICE` | `cpu` \| `cuda` \| `auto` (auto = try GPU, fall back to CPU) | `cpu` | +| `WHISPER_MODEL` | Model size: `tiny`, `base`, `small`, `medium`, `large-v2`, `large-v3`, `large-v3-turbo` | `base` | +| `WHISPER_DEVICE` | `cpu` \| `cuda` | `cpu` | | `HF_TOKEN` | Hugging Face token for faster model downloads (optional; [create one](https://huggingface.co/settings/tokens)) | — | --- @@ -335,7 +335,7 @@ Browse: [model.lmstudio.ai](https://model.lmstudio.ai) | `ALLOWED_TELEGRAM_USER_ID` | Allowed Telegram User ID | `""` | | `VOICE_NOTE_ENABLED` | Enable voice note handling | `true` | | `WHISPER_MODEL` | Local Whisper model size | `base` | -| `WHISPER_DEVICE` | `cpu` \| `cuda` \| `auto` | `cpu` | +| `WHISPER_DEVICE` | `cpu` \| `cuda` | `cpu` | | `MESSAGING_RATE_LIMIT` | Messaging messages per window | `1` | | `MESSAGING_RATE_WINDOW` | Messaging window (seconds) | `1` | | `CLAUDE_WORKSPACE` | Directory for agent workspace | `./agent_workspace` | diff --git a/config/settings.py b/config/settings.py index 93e130b2..cd92f4c5 100644 --- a/config/settings.py +++ b/config/settings.py @@ -78,9 +78,9 @@ class Settings(BaseSettings): ) # Hugging Face token for faster model downloads (optional) hf_token: str = Field(default="", validation_alias="HF_TOKEN") - # Model size: "tiny" | "base" | "small" | "medium" | "large-v2" + # Model size: "tiny" | "base" | "small" | "medium" | "large-v2" | "large-v3" | "large-v3-turbo" whisper_model: str = Field(default="base", validation_alias="WHISPER_MODEL") - # Device: "cpu" | "cuda" | "auto" (auto = try cuda, fall back to cpu) + # Device: "cpu" | "cuda" whisper_device: str = Field(default="cpu", validation_alias="WHISPER_DEVICE") # ==================== Bot Wrapper Config ==================== @@ -115,6 +115,13 @@ class Settings(BaseSettings): return None return v + @field_validator("whisper_device") + @classmethod + def validate_whisper_device(cls, v: str) -> str: + if v not in ("cpu", "cuda"): + raise ValueError(f"whisper_device must be 'cpu' or 'cuda', got {v!r}") + return v + model_config = SettingsConfigDict( env_file=".env", env_file_encoding="utf-8", diff --git a/messaging/transcription.py b/messaging/transcription.py index 5ca5c556..3e21f082 100644 --- a/messaging/transcription.py +++ b/messaging/transcription.py @@ -15,8 +15,6 @@ MAX_AUDIO_SIZE_BYTES = 25 * 1024 * 1024 # Lazy-loaded models: (model_name, device) -> model _model_cache: dict[tuple[str, str], Any] = {} -# Models for which CUDA failed at inference; skip cuda on subsequent requests -_cuda_failed_models: set[str] = set() class _WhisperModelLike(Protocol): @@ -27,11 +25,10 @@ class _WhisperModelLike(Protocol): def _get_local_model(whisper_model: str, device: str) -> _WhisperModelLike: """Lazy-load faster-whisper model. Raises ImportError if not installed.""" - global _model_cache, _cuda_failed_models - resolved = device if device in ("cpu", "cuda") else "auto" - if resolved in ("cuda", "auto") and whisper_model in _cuda_failed_models: - resolved = "cpu" - cache_key = (whisper_model, resolved) + global _model_cache + if device not in ("cpu", "cuda"): + raise ValueError(f"whisper_device must be 'cpu' or 'cuda', got {device!r}") + cache_key = (whisper_model, device) if cache_key not in _model_cache: try: from config.settings import get_settings @@ -43,19 +40,12 @@ def _get_local_model(whisper_model: str, device: str) -> _WhisperModelLike: faster_whisper = importlib.import_module("faster_whisper") WhisperModel = faster_whisper.WhisperModel - if resolved == "auto": - try: - _model_cache[cache_key] = WhisperModel(whisper_model, device="cuda") - except RuntimeError: - _model_cache[cache_key] = WhisperModel( - whisper_model, device="cpu", compute_type="float32" - ) - elif resolved == "cpu": + if device == "cpu": _model_cache[cache_key] = WhisperModel( whisper_model, device="cpu", compute_type="float32" ) else: - _model_cache[cache_key] = WhisperModel(whisper_model, device=resolved) + _model_cache[cache_key] = WhisperModel(whisper_model, device=device) except ImportError as e: raise ImportError( "Voice notes require the voice extra. Install with: uv sync --extra voice" @@ -76,8 +66,8 @@ def transcribe_audio( Args: file_path: Path to audio file (OGG, MP3, MP4, WAV, M4A supported) mime_type: MIME type of the audio (e.g. "audio/ogg") - whisper_model: Model size: "tiny", "base", "small", "medium", "large-v2" - whisper_device: "cpu" | "cuda" | "auto" (auto = try GPU, fall back to CPU) + whisper_model: Model size: "tiny", "base", "small", "medium", "large-v2", "large-v3", "large-v3-turbo" + whisper_device: "cpu" | "cuda" Returns: Transcribed text @@ -100,23 +90,9 @@ def transcribe_audio( def _transcribe_local(file_path: Path, whisper_model: str, whisper_device: str) -> str: - """Transcribe using local faster-whisper.""" + """Transcribe using local faster-whisper. Fails fast on device errors (no fallback).""" model: _WhisperModelLike = _get_local_model(whisper_model, whisper_device) - try: - segments, _info = model.transcribe(str(file_path), beam_size=5) - except RuntimeError as e: - err_lower = str(e).lower() - if "cublas" in err_lower or "cuda" in err_lower: - # CUDA deferred load failed at inference; remember and fall back to CPU - global _model_cache, _cuda_failed_models - _cuda_failed_models.add(whisper_model) - for key in list(_model_cache): - if key[0] == whisper_model: - del _model_cache[key] - model = _get_local_model(whisper_model, "cpu") - segments, _info = model.transcribe(str(file_path), beam_size=5) - else: - raise + segments, _info = model.transcribe(str(file_path), beam_size=5) parts = [s.text for s in segments if s.text] result = " ".join(parts).strip() logger.debug(f"Local transcription: {len(result)} chars") diff --git a/tests/config/test_config.py b/tests/config/test_config.py index 269ab724..758d4dd5 100644 --- a/tests/config/test_config.py +++ b/tests/config/test_config.py @@ -303,3 +303,20 @@ class TestSettingsOptionalStr: monkeypatch.setenv("MESSAGING_PLATFORM", "discord") s = Settings() assert s.messaging_platform == "discord" + + def test_whisper_device_auto_rejected(self, monkeypatch): + """WHISPER_DEVICE=auto raises ValidationError (auto removed).""" + from config.settings import Settings + + monkeypatch.setenv("WHISPER_DEVICE", "auto") + with pytest.raises(ValidationError, match="whisper_device"): + Settings() + + @pytest.mark.parametrize("device", ["cpu", "cuda"]) + def test_whisper_device_valid(self, monkeypatch, device): + """Valid whisper_device values are accepted.""" + from config.settings import Settings + + monkeypatch.setenv("WHISPER_DEVICE", device) + s = Settings() + assert s.whisper_device == device diff --git a/tests/messaging/test_transcription.py b/tests/messaging/test_transcription.py index e58237fd..2cb4624b 100644 --- a/tests/messaging/test_transcription.py +++ b/tests/messaging/test_transcription.py @@ -73,6 +73,18 @@ def test_transcribe_local_empty_segments_returns_no_speech(): path.unlink(missing_ok=True) +def test_transcribe_invalid_device_raises(): + """Invalid whisper_device raises ValueError.""" + with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: + f.write(b"fake ogg") + path = Path(f.name) + try: + with pytest.raises(ValueError, match="whisper_device must be 'cpu' or 'cuda'"): + transcribe_audio(path, "audio/ogg", whisper_device="auto") + finally: + path.unlink(missing_ok=True) + + def test_transcribe_local_import_error_raises(): """Local backend when faster-whisper not installed raises ImportError.""" with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: