Merge pull request #40 from Alishahryar1/cursor/voice-note-transcribe-fallbacks-adb6

Voice note transcribe fallbacks
This commit is contained in:
Ali Khokhar
2026-02-18 05:39:39 -08:00
committed by GitHub
6 changed files with 52 additions and 40 deletions
+1 -1
View File
@@ -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
+3 -3
View File
@@ -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` |
+9 -2
View File
@@ -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",
+10 -34
View File
@@ -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")
+17
View File
@@ -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
+12
View File
@@ -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: