feat: pluggable backends + OpenVINO для ускорения на x86 CPU

Добавлена pluggable-архитектура бэкендов транскрипции и OpenVINO
как второй движок для ускорения на Intel/AMD CPU в 3-6 раз.

- Backend Protocol (structural typing) + реестр с lazy imports
- FasterWhisperBackend (CUDA/CPU) — рефакторинг без изменения поведения
- OpenVINOBackend — openvino-genai WhisperPipeline, предквантизированные модели
- Auto-detect: CUDA → OpenVINO → CPU
- Cross-backend fallback с сохранением состояния в батч-режиме
- Тесты на 3 CPU: Intel Ultra 7, AMD Ryzen 7, Intel i7 (WSL2)

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-03-22 12:16:35 +03:00
co-authored by Claude Opus 4.6
19 changed files with 1629 additions and 639 deletions
+60 -20
View File
@@ -8,7 +8,7 @@ transcribe meeting.mp4
``` ```
- **Полностью локально** — данные не покидают машину - **Полностью локально** — данные не покидают машину
- **Авто-GPU** — автоматически использует NVIDIA CUDA, если доступен - **Авто-ускорение** — NVIDIA CUDA, OpenVINO (Intel/AMD CPU) или CPU fallback
- **Батч-режим** — обработка нескольких файлов за один вызов - **Батч-режим** — обработка нескольких файлов за один вызов
- **Markdown с таймкодами** — удобен для суммаризации ИИ - **Markdown с таймкодами** — удобен для суммаризации ИИ
- **Аудио и видео** — mp3, wav, mp4, mkv и [другие форматы](#поддерживаемые-форматы) - **Аудио и видео** — mp3, wav, mp4, mkv и [другие форматы](#поддерживаемые-форматы)
@@ -30,12 +30,12 @@ powershell -ExecutionPolicy ByPass -c "irm https://astral.sh/uv/install.ps1 | ie
uv tool install git+https://github.com/dementev-dev/local-transcriber uv tool install git+https://github.com/dementev-dev/local-transcriber
``` ```
**3. (Опционально) GPU-ускорение:** **3. Ускорение (ставится автоматически):**
Если есть NVIDIA GPU — транскрипция будет в 5–10× быстрее. Требуется **CUDA 12** (ctranslate2 4.7 не совместим с CUDA 11 и 13). - **OpenVINO** (Intel/AMD x86 CPU): ставится автоматически на Linux и Windows — ускорение в 2-4 раза
- **NVIDIA CUDA** (GPU): если есть GPU — транскрипция в 5-10× быстрее
- **Windows**: `winget install -e --id Nvidia.CUDA --version 12.9` (от администратора), перезапустить терминал - **Windows**: `winget install -e --id Nvidia.CUDA --version 12.9` (от администратора), перезапустить терминал
- **Linux / WSL2**: работает из коробки (нужен только драйвер: `nvidia-smi`) - **Linux / WSL2**: работает из коробки (нужен только драйвер: `nvidia-smi`)
**4. Готово:** **4. Готово:**
@@ -63,6 +63,41 @@ uv tool install --force git+https://github.com/dementev-dev/local-transcriber
uv tool uninstall local-transcriber uv tool uninstall local-transcriber
``` ```
**Очистка моделей:**
Модели кешируются в `~/.cache/huggingface/hub/` и могут занимать несколько гигабайт.
На Windows без Developer Mode файлы копируются без симлинков — место удваивается.
```bash
# Linux / macOS — посмотреть размер кеша
du -sh ~/.cache/huggingface/hub/models--*
# Удалить все скачанные модели
rm -rf ~/.cache/huggingface/hub/models--Systran--faster-whisper-*
rm -rf ~/.cache/huggingface/hub/models--OpenVINO--whisper-*
```
```powershell
# Windows
dir "$env:USERPROFILE\.cache\huggingface\hub\models--*"
# Удалить все скачанные модели
Remove-Item -Recurse "$env:USERPROFILE\.cache\huggingface\hub\models--Systran--faster-whisper-*"
Remove-Item -Recurse "$env:USERPROFILE\.cache\huggingface\hub\models--OpenVINO--whisper-*"
```
При следующем запуске нужная модель скачается заново.
<details>
<summary>Windows: ошибка WinError 1314 при первом запуске</summary>
HuggingFace Hub использует симлинки для экономии места. На Windows без Developer Mode первая загрузка модели может упасть с ошибкой `WinError 1314`. Повторный запуск команды обычно помогает — HF Hub переключается на копирование файлов.
Чтобы избежать проблемы и сэкономить место, включите Developer Mode:
[Инструкция Microsoft](https://docs.microsoft.com/en-us/windows/apps/get-started/enable-your-device-for-development)
</details>
## Использование ## Использование
```bash ```bash
@@ -106,8 +141,8 @@ transcribe *.mp4 --force
| `--model` | `-m` | `medium` | Модель Whisper | | `--model` | `-m` | `medium` | Модель Whisper |
| `--language` | `-l` | `ru` | Язык (ru, en, auto и др.) | | `--language` | `-l` | `ru` | Язык (ru, en, auto и др.) |
| `--output` | `-o` | `<файл>-transcript.md` | Путь к выходному файлу | | `--output` | `-o` | `<файл>-transcript.md` | Путь к выходному файлу |
| `--device` | `-d` | `auto` | Устройство (auto, cpu, cuda) | | `--device` | `-d` | `auto` | Устройство (auto, cpu, cuda, openvino) |
| `--compute-type` | — | float16 (GPU) / float32 (CPU) | Тип вычислений | | `--compute-type` | — | float16 (CUDA) / int8 (OpenVINO) / float32 (CPU) | Тип вычислений |
| `--force` | `-f` | — | Перезаписать существующие транскрипты | | `--force` | `-f` | — | Перезаписать существующие транскрипты |
| `--verbose` | `-v` | — | Подробный вывод | | `--verbose` | `-v` | — | Подробный вывод |
@@ -116,6 +151,7 @@ transcribe *.mp4 --force
| | Linux / WSL2 | macOS | Windows | | | Linux / WSL2 | macOS | Windows |
|---|---|---|---| |---|---|---|---|
| CPU | ✅ | ✅ | ✅ | | CPU | ✅ | ✅ | ✅ |
| OpenVINO (x86 CPU) | ✅ авто | — | ✅ авто |
| GPU (NVIDIA) | ✅ авто | — | ✅ (нужен CUDA 12) | | GPU (NVIDIA) | ✅ авто | — | ✅ (нужен CUDA 12) |
<details> <details>
@@ -166,11 +202,11 @@ language = "en"
Дефолты зависят от устройства: Дефолты зависят от устройства:
| Параметр | GPU (CUDA) | CPU | | Параметр | CUDA | OpenVINO | CPU |
|----------|-----------|-----| |----------|------|----------|-----|
| model | medium | medium | | model | medium | medium | medium |
| compute_type | float16 | float32 | | compute_type | float16 | int8 | float32 |
| language | ru | ru | | language | ru | ru | ru |
## Модели и GPU ## Модели и GPU
@@ -195,19 +231,23 @@ language = "en"
<details> <details>
<summary>Типы квантизации (--compute-type)</summary> <summary>Типы квантизации (--compute-type)</summary>
| Тип | Устройство | VRAM/RAM | Качество | Когда использовать | | Тип | Бэкенд | VRAM/RAM | Качество | Когда использовать |
|-----|-----------|----------|----------|--------------------| |-----|--------|----------|----------|--------------------|
| `float16` | GPU | ~4.5-5 GB | Отлично | **По умолчанию для GPU** | | `float16` | CUDA | ~4.5-5 GB | Отлично | **По умолчанию для CUDA** |
| `int8_float16` | GPU | ~4.7 GB | Отлично | GPU от 6 GB, альтернатива float16 | | `int8_float16` | CUDA | ~4.7 GB | Отлично | GPU от 6 GB, альтернатива float16 |
| `int8` | GPU/CPU | Низкое | Хорошо, но бывают галлюцинации | GPU от 4 GB, CPU | | `int8` | CUDA / OpenVINO | Низкое | Хорошо, но бывают галлюцинации | **По умолчанию для OpenVINO** |
| `fp16` | OpenVINO | Низкое | Отлично | OpenVINO large-v3 (выбирается автоматически) |
| `float32` | CPU | Среднее | Отлично | **По умолчанию для CPU** | | `float32` | CPU | Среднее | Отлично | **По умолчанию для CPU** |
**Важно:** `int8` на длинных записях может давать галлюцинации (повтор фраз, потеря контента). **Важно:** `int8` на длинных записях может давать галлюцинации (повтор фраз, потеря контента).
`float16` и `float32` значительно стабильнее на записях >20 минут. `float16`/`fp16` и `float32` значительно стабильнее на записях >20 минут.
> Для OpenVINO `--compute-type` выбирает предквантизированную модель (int8 или fp16),
> а не runtime-параметр. Для `large-v3` по умолчанию выбирается `fp16`.
</details> </details>
Подробнее: бенчмарки, совместимость GPU, результаты тестирования — [docs/gpu.md](docs/gpu.md). Подробнее: бенчмарки, OpenVINO, совместимость GPU, результаты тестирования — [docs/gpu.md](docs/gpu.md).
<details> <details>
<summary>Формат вывода</summary> <summary>Формат вывода</summary>
+96
View File
@@ -0,0 +1,96 @@
# ADR-003: Pluggable backends и OpenVINO
**Статус**: Принято
**Дата**: 2026-03-21
## Контекст
На CPU (faster-whisper/CTranslate2) транскрипция работает медленно (~1.5x реалтайм для medium).
CUDA доступна на малом проценте машин (ноутбуки с NVIDIA GPU), на офисных ПК её нет.
OpenVINO ускоряет inference на x86 CPU (Intel и AMD) в 2-4 раза. Для его поддержки
нужен второй движок транскрипции, а архитектура должна позволять добавлять новые
бэкенды (CoreML для Mac, AMD XDNA NPU) без переписывания существующего кода.
## Решение
### Backend Protocol (structural typing)
Минимальный интерфейс в `backends/base.py`:
```python
class Backend(Protocol):
def ensure_model_available(self, model_name, compute_type, on_status) -> str: ...
def create_model(self, model_path, device, compute_type) -> Any: ...
def transcribe(self, model, file_path, language, on_segment, on_status) -> TranscribeResult: ...
```
Protocol вместо ABC — бэкенды не наследуются, достаточно реализовать методы.
Соответствует стилю проекта (наследование нигде не используется).
### Ленивые импорты
Бэкенды импортируются только при выборе — `get_backend(device)` делает import внутри.
Импорт faster-whisper запускает CUDA bootstrap (~1ms), импорт openvino-genai загружает ~50MB
shared libraries. Ни то, ни другое не должно происходить, если бэкенд не выбран.
### Device как селектор бэкенда
Вместо отдельного `--backend` флага устройство само определяет бэкенд:
- `cuda`, `cpu` → FasterWhisperBackend
- `openvino` → OpenVINOBackend
- `auto` → CUDA (nvidia-smi) → OpenVINO (import check + x86) → CPU
### load_model() — единственный владелец pipeline
`load_model()` выполняет ensure_model_available + create_model в одном вызове.
CLI не вызывает ensure_model_available отдельно — это убирает двойной resolution
и гарантирует, что модель скачивается для правильного бэкенда.
### Cross-backend fallback
Fallback живёт в `transcriber.py` (оркестратор), не в бэкендах:
- CUDA ошибка → CPU (FasterWhisper)
- OpenVINO ошибка → CPU (FasterWhisper)
- `strict_device=True` (явный `--device`) → ошибка без fallback
При fallback в батч-режиме обновляются model, backend, model_path и actual_device
через TranscribeFileResult — следующий файл использует правильный бэкенд.
### Аудио для OpenVINO
OpenVINO GenAI WhisperPipeline принимает raw PCM float массив, не путь к файлу.
Используем `faster_whisper.decode_audio()` (PyAV) → `.tolist()``pipe.generate()`.
Системный ffmpeg не требуется — PyAV бандлит FFmpeg внутри wheel.
### compute_type для OpenVINO
OpenVINO модели предквантизированы (int8/fp16), compute_type определяет какую модель
скачать. Контракт:
- Явный `--compute-type` или значение из конфига — уважается всегда
- Из дефолтов: для large-v3 автоматически выбирается fp16 (стабильнее по качеству)
- Несуществующая пара (model + compute_type) при явном выборе → ошибка
### Обе зависимости по умолчанию
faster-whisper (~37MB) и openvino-genai (~69MB) ставятся вместе — суммарно ~106MB,
приемлемо. Модели скачиваются только для активного бэкенда. CUDA (nvidia-cublas-cu12,
~554MB) остаётся conditional (Linux x86_64). OpenVINO — conditional (x86_64/AMD64, не macOS).
## Последствия
- Обратная совместимость: `transcribe()` сохранён; `load_model()` изменил сигнатуру (возвращает 4-tuple вместо 2-tuple, добавлен `compute_type_explicit`)
- Новый бэкенд добавляется одним файлом в `backends/` + регистрацией в `__init__.py`
- Модели скачиваются по запросу — CUDA пользователь не качает OpenVINO модели, и наоборот
- ARM и macOS: OpenVINO не ставится (platform markers), работает CPU через faster-whisper
## Отклонённые альтернативы
| Альтернатива | Почему отклонена |
|---|---|
| OpenVINO как optional extra (`pip install .[openvino]`) | Теряется zero-config UX; пользователь должен знать про extras |
| whisper.cpp (pywhispercpp) | Другой движок, больший объём интеграции; OpenVINO GenAI проще |
| Единый бэкенд с OpenVINO для всего | CTranslate2 лучше оптимизирован для CUDA; OpenVINO — для CPU |
| ABC вместо Protocol | Наследование не используется в проекте; Protocol проще |
| librosa для загрузки аудио в OpenVINO | Лишняя зависимость; для видеоконтейнеров ненадёжна без системного ffmpeg |
| `--backend` как отдельный флаг | Усложняет CLI; device уже однозначно определяет бэкенд |
+81 -16
View File
@@ -1,10 +1,78 @@
# GPU и CUDA # Ускорение транскрипции
## Режимы `--device` ## Режимы `--device`
- `auto` (по умолчанию) — выберет GPU если `nvidia-smi` доступен, иначе CPU - `auto` (по умолчанию) — CUDA → OpenVINO → CPU (первый доступный)
- `cuda` — строго GPU, ошибка если недоступен (без silent fallback) - `cuda` — строго NVIDIA GPU, ошибка если недоступен
- `cpu` — строго CPU - `openvino` — OpenVINO на CPU (ускорение 2-4x на x86)
- `cpu` — строго CPU (faster-whisper/CTranslate2)
## Какой бэкенд на каком оборудовании
| Оборудование | Рекомендуемый `--device` | Бэкенд | Ожидаемая скорость |
|---|---|---|---|
| NVIDIA GPU (6+ GB VRAM) | `auto` / `cuda` | faster-whisper (CTranslate2) | 7-19x реалтайм |
| Intel/AMD x86 CPU | `auto` / `openvino` | OpenVINO GenAI | 3-6x реалтайм* |
| Любой CPU (fallback) | `cpu` | faster-whisper (CTranslate2) | ~1.5x реалтайм |
| Apple Silicon (macOS) | `cpu` | faster-whisper (CTranslate2) | ~2x реалтайм |
\* По результатам тестирования на Intel и AMD CPU. Реальная скорость зависит от CPU и модели.
## OpenVINO
OpenVINO ускоряет inference на x86 процессорах (Intel и AMD) через оптимизированные инструкции
(AVX2, AVX-512, VNNI, AMX). Ставится автоматически на Linux и Windows (x86_64/AMD64).
- **Модели**: предконвертированные из [HuggingFace](https://huggingface.co/OpenVINO) (int8/fp16)
- **Дефолт**: `medium` + `int8` (для `large-v3` автоматически выбирается `fp16`)
- **Аудиодекодирование**: через PyAV (бандлит FFmpeg), системный ffmpeg не нужен
### Доступные OpenVINO модели
| Модель | int8 | fp16 |
|--------|------|------|
| tiny | OpenVINO/whisper-tiny-int8-ov | — |
| base | — | OpenVINO/whisper-base-fp16-ov |
| small | OpenVINO/whisper-small-int8-ov | — |
| medium | OpenVINO/whisper-medium-int8-ov | — |
| large-v3 | OpenVINO/whisper-large-v3-int8-ov | OpenVINO/whisper-large-v3-fp16-ov |
### Результаты тестирования OpenVINO
Реальные записи рабочих созвонов (русский, техтермины: SQL, PostgreSQL, LDAP, DLP и др.).
**Скорость (эталонный файл 16 мин, OpenVINO, medium int8):**
| CPU | medium int8 | large-v3 fp16 | CPU float32 (baseline) |
|---|---|---|---|
| Intel Ultra 7 255H | **122с** | **411с** | — |
| AMD Ryzen 7 8845H | 185с | 416с | 734с |
| Intel i7 (WSL2) | 171-205с | — | 658с |
**Ускорение vs CPU float32:** **3-6x** в зависимости от CPU.
**OpenVINO small int8 (Intel i7 WSL2):**
| Файл | small int8 | medium int8 |
|---|---|---|
| 16 мин | 93с | 171с |
| 42 мин | 153с (~16x реалтайм) | 413с |
**Качество (сравнение на одном файле, 16 мин, OpenVINO int8/fp16, Intel CPU):**
| | small int8 | medium int8 | large-v3 fp16 |
|---|---|---|---|
| Время (16 мин) | **93с** | 171с | 416с |
| Время (42 мин) | **153с** | 413с | — |
| Ключевые слова | искажения ("бокап", "рецензия") | единичные ляпы | корректно |
| Пунктуация | слабая | базовая | хорошая |
| Галлюцинации | нет | нет | нет |
### Рекомендации по выбору модели
- **small** — для быстрого сканирования большого объёма видео по маске (`*.mp4`). Ошибки в отдельных словах; для обработки ИИ (МОМ, конспект) рискованно — "рецензия" вместо "лицензия" может исказить смысл.
- **medium** — для повседневного использования и обработки ИИ. Ключевые термины верные, единичные ляпы не влияют на смысл конспекта. Оптимальный баланс скорости и качества.
- **large-v3** — для важных записей, где нужна дословная точность. Лучшая пунктуация и связность. На OpenVINO (416с) быстрее, чем medium на чистом CPU (734с) — лучшее качество при выше скорости.
## Настройка по платформам ## Настройка по платформам
@@ -17,12 +85,10 @@
### Windows ### Windows
Нужен системный CUDA toolkit: Нужен системный **CUDA 12** (ctranslate2 4.7 не совместим с CUDA 11 и 13):
```bash ```bash
choco install cuda winget install -e --id Nvidia.CUDA --version 12.9 # требует запуска от имени администратора
# или
winget install -e --id Nvidia.CUDA # требует запуска от имени администратора
``` ```
После установки перезапустите терминал. После установки перезапустите терминал.
@@ -43,10 +109,11 @@ winget install -e --id Nvidia.CUDA # требует запуска от име
|-------------|-------------|-------------|---------------| |-------------|-------------|-------------|---------------|
| GPU + medium float16 | ~35с | ~133с | ~19x реалтайм | | GPU + medium float16 | ~35с | ~133с | ~19x реалтайм |
| GPU + large-v3 float16 | ~90с | ~350с | ~7x реалтайм | | GPU + large-v3 float16 | ~90с | ~350с | ~7x реалтайм |
| CPU + medium float32 | 613с (10 мин) | ~26 мин* | ~1.5x реалтайм | | **OpenVINO + small int8** | **93с** | **153с** | **~10-16x реалтайм** |
| CPU + large-v3 int8 | 839с (14 мин) | ~37 мин* | ~1:1 реалтайм | | **OpenVINO + medium int8** | **171-205с** | **413с** | **~4-6x реалтайм** |
| **OpenVINO + large-v3 fp16** | **416с** | — | **~2.3x реалтайм** |
*Оценка на основе пропорции. | CPU + medium float32 | 658с (11 мин) | ~26 мин | ~1.5x реалтайм |
| CPU + large-v3 int8 | 839с (14 мин) | ~37 мин | ~1:1 реалтайм |
## Результаты тестирования качества ## Результаты тестирования качества
@@ -77,11 +144,9 @@ SQL, PostgreSQL, Greenplum, Airflow, ClickHouse, Docker, CDR, GTP, MAP).
### Windows: ошибка при загрузке модели на GPU ### Windows: ошибка при загрузке модели на GPU
GPU на Windows требует CUDA toolkit (включает cuBLAS). Установите: GPU на Windows требует **CUDA 12** (ctranslate2 4.7 не совместим с CUDA 11 и 13). Установите:
```bash ```bash
choco install cuda winget install -e --id Nvidia.CUDA --version 12.9 # требует запуска от имени администратора
# или
winget install -e --id Nvidia.CUDA
``` ```
После установки перезапустите терминал. После установки перезапустите терминал.
+1
View File
@@ -10,6 +10,7 @@ dependencies = [
"faster-whisper>=1.2.1", "faster-whisper>=1.2.1",
"socksio>=1.0.0", "socksio>=1.0.0",
"nvidia-cublas-cu12>=12.4; sys_platform == 'linux' and platform_machine == 'x86_64'", "nvidia-cublas-cu12>=12.4; sys_platform == 'linux' and platform_machine == 'x86_64'",
"openvino-genai>=2025.0; sys_platform != 'darwin' and (platform_machine == 'x86_64' or platform_machine == 'AMD64')",
"tomli>=2.0; python_version < '3.11'", "tomli>=2.0; python_version < '3.11'",
] ]
@@ -0,0 +1,29 @@
"""Реестр бэкендов транскрипции и выбор бэкенда по устройству."""
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from .base import Backend
def get_backend(device: str, *, compute_type_explicit: bool = True) -> Backend:
"""Возвращает экземпляр бэкенда для указанного устройства.
Импорты ленивые — бэкенд загружается только при запросе.
compute_type_explicit: False если compute_type пришёл из дефолтов (влияет на fallback).
"""
if device == "openvino":
try:
from .openvino import OpenVINOBackend
except ImportError:
raise ValueError(
"OpenVINO бэкенд недоступен. Установите: pip install openvino-genai"
) from None
return OpenVINOBackend(compute_type_explicit=compute_type_explicit)
# cuda, cpu и всё остальное → faster-whisper
from .faster_whisper import FasterWhisperBackend
return FasterWhisperBackend()
+46
View File
@@ -0,0 +1,46 @@
"""Протокол бэкенда транскрипции."""
from __future__ import annotations
from collections.abc import Callable
from pathlib import Path
from typing import Any, Protocol
from local_transcriber.types import Segment, TranscribeResult
class Backend(Protocol):
"""Минимальный интерфейс бэкенда транскрипции.
Бэкенды реализуют этот протокол (structural typing) —
наследование не требуется.
"""
def ensure_model_available(
self,
model_name: str,
compute_type: str,
on_status: Callable[[str], None] | None = None,
) -> str:
"""Гарантирует наличие модели, возвращает путь к файлам."""
...
def create_model(
self,
model_path: str,
device: str,
compute_type: str,
) -> Any:
"""Создаёт модель. Возвращает backend-специфичный объект."""
...
def transcribe(
self,
model: Any,
file_path: Path,
language: str | None,
on_segment: Callable[[Segment], None] | None = None,
on_status: Callable[[str], None] | None = None,
) -> TranscribeResult:
"""Транскрибирует файл, возвращает результат."""
...
@@ -0,0 +1,185 @@
"""Бэкенд транскрипции на основе faster-whisper (CTranslate2)."""
from __future__ import annotations
import gc
import io
import warnings
from collections.abc import Callable
from pathlib import Path
from typing import Any
# CUDA bootstrap — должен быть ДО импорта faster_whisper / ctranslate2
from local_transcriber._cuda_bootstrap import ensure_cublas_loadable
ensure_cublas_loadable()
from faster_whisper import WhisperModel # noqa: E402
from huggingface_hub import snapshot_download # noqa: E402
from huggingface_hub.errors import LocalEntryNotFoundError # noqa: E402
from local_transcriber.types import Segment, TranscribeResult # noqa: E402
MODEL_REPOS = {
"tiny": "Systran/faster-whisper-tiny",
"base": "Systran/faster-whisper-base",
"small": "Systran/faster-whisper-small",
"medium": "Systran/faster-whisper-medium",
"large-v3": "Systran/faster-whisper-large-v3",
}
MODEL_ALLOW_PATTERNS = [
"config.json",
"preprocessor_config.json",
"model.bin",
"tokenizer.json",
"vocabulary.*",
]
MODEL_REQUIRED_FILES = [
"config.json",
"model.bin",
"tokenizer.json",
]
class FasterWhisperBackend:
"""Бэкенд транскрипции через faster-whisper (CTranslate2)."""
def __init__(self):
self.actual_compute_type: str | None = None
def ensure_model_available(
self,
model_name: str,
compute_type: str,
on_status: Callable[[str], None] | None = None,
) -> str:
"""Резолвит alias модели в repo_id и гарантирует наличие файлов."""
self.actual_compute_type = compute_type
local_path = Path(model_name).expanduser()
if local_path.is_dir():
_validate_model_dir(local_path)
return str(local_path)
repo_id = _resolve_model_repo(model_name)
try:
_notify(on_status, f"Проверяю кэш модели {model_name}...")
cached_path = Path(_snapshot_download(repo_id, local_files_only=True))
_validate_model_dir(cached_path)
return str(cached_path)
except LocalEntryNotFoundError:
pass
except ValueError:
_notify(on_status, f"Кэш модели {model_name} неполный, докачиваю...")
_notify(on_status, f"Скачиваю модель {model_name} из Hugging Face...")
downloaded_path = Path(_snapshot_download(repo_id, local_files_only=False))
_validate_model_dir(downloaded_path)
return str(downloaded_path)
def create_model(
self,
model_path: str,
device: str,
compute_type: str,
) -> Any:
"""Создаёт WhisperModel."""
try:
return WhisperModel(model_path, device=device, compute_type=compute_type)
except ImportError as exc:
if _is_missing_socksio_error(exc):
raise RuntimeError(
"Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, "
"нужная для загрузки модели из Hugging Face через proxy. "
"Обновите окружение: `uv sync`."
) from exc
raise
def transcribe(
self,
model: Any,
file_path: Path,
language: str | None,
on_segment: Callable[[Segment], None] | None = None,
on_status: Callable[[str], None] | None = None,
) -> TranscribeResult:
"""Транскрибирует файл через faster-whisper."""
segment_generator, info = model.transcribe(
str(file_path), language=language,
)
total_duration = info.duration
segments: list[Segment] = []
for raw_seg in segment_generator:
seg = Segment(start=raw_seg.start, end=raw_seg.end, text=raw_seg.text)
if on_segment is not None:
on_segment(seg)
segments.append(seg)
_notify(
on_status,
f"Транскрибирую... {_fmt_time(seg.end)} / {_fmt_time(total_duration)}"
f" [{len(segments)} сегм.]",
)
return TranscribeResult(
segments=segments,
language=info.language,
language_probability=info.language_probability,
duration=info.duration,
device_used="", # оркестратор проставит actual_device
)
def _notify(on_status: Callable[[str], None] | None, message: str) -> None:
if on_status is not None:
on_status(message)
def _fmt_time(seconds: float) -> str:
m, s = divmod(int(seconds), 60)
h, m = divmod(m, 60)
return f"{h}:{m:02d}:{s:02d}" if h else f"{m:02d}:{s:02d}"
def _resolve_model_repo(model_name: str) -> str:
if "/" in model_name:
return model_name
repo_id = MODEL_REPOS.get(model_name)
if repo_id is None:
expected = ", ".join(MODEL_REPOS)
raise ValueError(f"Неподдерживаемая модель '{model_name}'. Ожидалось одно из: {expected}")
return repo_id
def _snapshot_download(repo_id: str, local_files_only: bool) -> str:
try:
return snapshot_download(
repo_id,
local_files_only=local_files_only,
allow_patterns=MODEL_ALLOW_PATTERNS,
)
except ImportError as exc:
if _is_missing_socksio_error(exc):
raise RuntimeError(
"Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, "
"нужная для загрузки модели из Hugging Face через proxy. "
"Обновите окружение: `uv sync`."
) from exc
raise
def _validate_model_dir(model_dir: Path) -> None:
missing = [
filename for filename in MODEL_REQUIRED_FILES if not (model_dir / filename).exists()
]
if not any(model_dir.glob("vocabulary.*")):
missing.append("vocabulary.*")
if missing:
missing_str = ", ".join(missing)
raise ValueError(f"Неполная локальная модель в '{model_dir}': отсутствуют {missing_str}")
def _is_missing_socksio_error(exc: BaseException) -> bool:
msg = str(exc).lower()
return "socks proxy" in msg and "socksio" in msg
+223
View File
@@ -0,0 +1,223 @@
"""Бэкенд транскрипции на основе OpenVINO GenAI."""
from __future__ import annotations
import threading
import time
import warnings
from collections.abc import Callable
from pathlib import Path
from typing import Any
from huggingface_hub import snapshot_download
from huggingface_hub.errors import LocalEntryNotFoundError
from local_transcriber.types import Segment, TranscribeResult
# (model_alias, compute_type) → HF repo
MODEL_REPOS: dict[tuple[str, str], str] = {
("tiny", "int8"): "OpenVINO/whisper-tiny-int8-ov",
("base", "fp16"): "OpenVINO/whisper-base-fp16-ov",
("small", "int8"): "OpenVINO/whisper-small-int8-ov",
("medium", "int8"): "OpenVINO/whisper-medium-int8-ov",
("large-v3", "int8"): "OpenVINO/whisper-large-v3-int8-ov",
("large-v3", "fp16"): "OpenVINO/whisper-large-v3-fp16-ov",
}
# Fallback: если точная пара не найдена, пробуем альтернативный compute_type
_COMPUTE_TYPE_FALLBACKS: dict[str, list[str]] = {
"float32": ["fp16", "int8"],
"float16": ["fp16", "int8"],
"fp16": ["fp16", "int8"],
"int8": ["int8", "fp16"],
}
# large-v3: при неявном compute_type предпочитаем fp16 (стабильнее по качеству)
_IMPLICIT_COMPUTE_TYPE_OVERRIDES: dict[str, str] = {
"large-v3": "fp16",
}
MODEL_REQUIRED_FILES = [
"openvino_encoder_model.xml",
"openvino_decoder_model.xml",
]
class OpenVINOBackend:
"""Бэкенд транскрипции через openvino-genai WhisperPipeline."""
def __init__(self, compute_type_explicit: bool = True):
"""compute_type_explicit=False означает, что compute_type пришёл из дефолтов."""
self._compute_type_explicit = compute_type_explicit
self.actual_compute_type: str | None = None
def ensure_model_available(
self,
model_name: str,
compute_type: str,
on_status: Callable[[str], None] | None = None,
) -> str:
"""Скачивает/находит OpenVINO модель нужной квантизации."""
repo_id, resolved_ct = self._resolve_repo(model_name, compute_type)
self.actual_compute_type = resolved_ct
try:
_notify(on_status, f"Проверяю кэш модели {model_name} (OpenVINO)...")
cached_path = Path(snapshot_download(repo_id, local_files_only=True))
_validate_model_dir(cached_path)
return str(cached_path)
except LocalEntryNotFoundError:
pass
except ValueError:
_notify(on_status, f"Кэш модели {model_name} неполный, докачиваю...")
_notify(on_status, f"Скачиваю модель {model_name} (OpenVINO) из Hugging Face...")
downloaded_path = Path(snapshot_download(repo_id, local_files_only=False))
_validate_model_dir(downloaded_path)
return str(downloaded_path)
def create_model(
self,
model_path: str,
device: str,
compute_type: str,
) -> Any:
"""Создаёт WhisperPipeline."""
import openvino_genai as ov_genai
return ov_genai.WhisperPipeline(model_path, "CPU")
def transcribe(
self,
model: Any,
file_path: Path,
language: str | None,
on_segment: Callable[[Segment], None] | None = None,
on_status: Callable[[str], None] | None = None,
) -> TranscribeResult:
"""Транскрибирует файл через OpenVINO GenAI."""
from faster_whisper import decode_audio
_notify(on_status, "Загружаю аудио...")
raw_speech = decode_audio(str(file_path), sampling_rate=16000)
duration = len(raw_speech) / 16000.0
kwargs: dict[str, Any] = {"return_timestamps": True}
if language:
kwargs["language"] = f"<|{language}|>"
dur_min = int(duration // 60)
duration_str = f"{dur_min} мин" if dur_min > 0 else f"{int(duration)} сек"
pcm_list = raw_speech.tolist()
result = _generate_with_progress(model, pcm_list, kwargs, duration_str, on_status)
segments: list[Segment] = []
if hasattr(result, "chunks") and result.chunks:
for chunk in result.chunks:
start = max(0.0, chunk.start_ts)
end = max(start, chunk.end_ts)
seg = Segment(
start=start,
end=end,
text=chunk.text,
)
if on_segment is not None:
on_segment(seg)
segments.append(seg)
_notify(
on_status,
f"Транскрибирую (OpenVINO)... [{len(segments)} сегм.]",
)
detected_language = language or "auto"
language_probability = 1.0 if language else 0.0
return TranscribeResult(
segments=segments,
language=detected_language,
language_probability=language_probability,
duration=duration,
device_used="", # оркестратор проставит
)
def _resolve_repo(self, model_name: str, compute_type: str) -> tuple[str, str]:
"""Находит HF repo для пары (model, compute_type) с fallback.
Возвращает (repo_id, actual_compute_type).
"""
# Для неявного compute_type: override для конкретных моделей
if not self._compute_type_explicit and model_name in _IMPLICIT_COMPUTE_TYPE_OVERRIDES:
compute_type = _IMPLICIT_COMPUTE_TYPE_OVERRIDES[model_name]
# Точное совпадение
repo = MODEL_REPOS.get((model_name, compute_type))
if repo:
return repo, compute_type
# Fallback только для неявного compute_type
if not self._compute_type_explicit:
fallbacks = _COMPUTE_TYPE_FALLBACKS.get(compute_type, [])
for fallback_ct in fallbacks:
repo = MODEL_REPOS.get((model_name, fallback_ct))
if repo:
return repo, fallback_ct
# Явный --compute-type с несуществующей парой → ошибка
available = [ct for (m, ct) in MODEL_REPOS if m == model_name]
if available:
raise ValueError(
f"Модель '{model_name}' недоступна с compute_type='{compute_type}' для OpenVINO. "
f"Доступные варианты: {', '.join(sorted(set(available)))}"
)
all_models = sorted({m for m, _ in MODEL_REPOS})
raise ValueError(
f"Модель '{model_name}' не найдена для OpenVINO. "
f"Доступные модели: {', '.join(all_models)}"
)
def _generate_with_progress(
model: Any,
pcm_list: list[float],
kwargs: dict[str, Any],
duration_str: str,
on_status: Callable[[str], None] | None,
) -> Any:
"""Запускает model.generate() в потоке, обновляя статус с elapsed time."""
result_box: list[Any] = [None]
error_box: list[BaseException | None] = [None]
def run() -> None:
try:
result_box[0] = model.generate(pcm_list, **kwargs)
except BaseException as exc:
error_box[0] = exc
thread = threading.Thread(target=run)
start = time.monotonic()
thread.start()
while thread.is_alive():
elapsed = int(time.monotonic() - start)
elapsed_str = f"{elapsed // 60:02d}:{elapsed % 60:02d}"
_notify(on_status, f"Транскрибирую {duration_str} аудио (OpenVINO)... прошло {elapsed_str}")
thread.join(timeout=1.0)
if error_box[0] is not None:
raise error_box[0]
return result_box[0]
def _notify(on_status: Callable[[str], None] | None, message: str) -> None:
if on_status is not None:
on_status(message)
def _validate_model_dir(model_dir: Path) -> None:
missing = [f for f in MODEL_REQUIRED_FILES if not (model_dir / f).exists()]
if missing:
raise ValueError(
f"Неполная OpenVINO модель в '{model_dir}': отсутствуют {', '.join(missing)}"
)
+46 -40
View File
@@ -14,9 +14,7 @@ from .transcriber import (
Segment, Segment,
_is_cuda_error, _is_cuda_error,
_transcribe_file, _transcribe_file,
ensure_model_available,
load_model, load_model,
transcribe,
) )
from .utils import ( from .utils import (
build_output_path, build_output_path,
@@ -31,6 +29,16 @@ app = typer.Typer()
console = Console(stderr=True) console = Console(stderr=True)
def _format_device_info(device_used: str) -> str:
"""Формирует строку устройства для шапки транскрипта."""
if device_used == "cuda":
gpu_name = get_gpu_name()
return f"CUDA ({gpu_name or 'Unknown GPU'})"
if device_used == "openvino":
return "OpenVINO (CPU)"
return "CPU"
@app.command() @app.command()
def main( def main(
files: list[Path] = typer.Argument(..., help="Пути к аудио/видеофайлам"), files: list[Path] = typer.Argument(..., help="Пути к аудио/видеофайлам"),
@@ -42,11 +50,12 @@ def main(
), ),
output: Path | None = typer.Option(None, "--output", "-o", help="Путь к выходному файлу"), output: Path | None = typer.Option(None, "--output", "-o", help="Путь к выходному файлу"),
device: str | None = typer.Option( device: str | None = typer.Option(
None, "--device", "-d", show_default=False, help="Устройство (auto|cpu|cuda) [по умолч.: auto]" None, "--device", "-d", show_default=False,
help="Устройство (auto|cpu|cuda|openvino) [по умолч.: auto]"
), ),
compute_type: str | None = typer.Option( compute_type: str | None = typer.Option(
None, "--compute-type", show_default=False, None, "--compute-type", show_default=False,
help="Тип вычислений [по умолч.: float16 (GPU) / float32 (CPU)]" help="Тип вычислений [по умолч.: float16 (CUDA) / int8 (OpenVINO) / float32 (CPU)]"
), ),
verbose: bool = typer.Option(False, "--verbose", "-v", help="Подробный вывод"), verbose: bool = typer.Option(False, "--verbose", "-v", help="Подробный вывод"),
force: bool = typer.Option(False, "--force", "-f", help="Перезаписать существующие транскрипты"), force: bool = typer.Option(False, "--force", "-f", help="Перезаписать существующие транскрипты"),
@@ -63,6 +72,8 @@ def main(
resolved_device = detect_device(defaults["device"]) resolved_device = detect_device(defaults["device"])
defaults = apply_device_defaults(defaults, resolved_device, cli_values, config) defaults = apply_device_defaults(defaults, resolved_device, cli_values, config)
ct_explicit = compute_type is not None or "compute_type" in config
expanded = expand_globs(files) expanded = expand_globs(files)
if not expanded: if not expanded:
console.print("Файлы не найдены.", style="red bold") console.print("Файлы не найдены.", style="red bold")
@@ -74,9 +85,9 @@ def main(
raise SystemExit(1) raise SystemExit(1)
if is_batch: if is_batch:
_run_batch(expanded, defaults, verbose, force) _run_batch(expanded, defaults, verbose, force, ct_explicit)
else: else:
_run_single(expanded[0], defaults, output, verbose) _run_single(expanded[0], defaults, output, verbose, ct_explicit)
except KeyboardInterrupt: except KeyboardInterrupt:
console.print("\nПрервано пользователем.", style="yellow") console.print("\nПрервано пользователем.", style="yellow")
raise SystemExit(130) raise SystemExit(130)
@@ -113,6 +124,7 @@ def _run_single(
defaults: dict[str, str], defaults: dict[str, str],
output: Path | None, output: Path | None,
verbose: bool, verbose: bool,
compute_type_explicit: bool = False,
) -> None: ) -> None:
"""Пайплайн одного файла: валидация → модель → транскрипция → запись.""" """Пайплайн одного файла: валидация → модель → транскрипция → запись."""
start = time.monotonic() start = time.monotonic()
@@ -120,35 +132,34 @@ def _run_single(
validated_file = validate_input_file(file) validated_file = validate_input_file(file)
requested_device = defaults["device"] requested_device = defaults["device"]
resolved_device = detect_device(requested_device) resolved_device = detect_device(requested_device)
# Если пользователь явно указал устройство — запрещаем fallback на CPU
strict = requested_device != "auto" strict = requested_device != "auto"
output_path = build_output_path(validated_file, output) output_path = build_output_path(validated_file, output)
console.print(f"Файл: [bold]{validated_file.name}[/bold]") console.print(f"Файл: [bold]{validated_file.name}[/bold]")
console.print(
f"Модель: [bold]{defaults['model']}[/bold] "
f"Устройство: [bold]{resolved_device}[/bold] "
f"Compute: [bold]{defaults['compute_type']}[/bold]"
)
model_path = ensure_model_available(
defaults["model"], on_status=lambda message: console.print(message)
)
def on_segment(seg: Segment) -> None: def on_segment(seg: Segment) -> None:
console.print(f" [{seg.start:.2f}s] {seg.text.strip()}") console.print(f" [{seg.start:.2f}s] {seg.text.strip()}")
model_obj, actual_device = load_model( model_obj, actual_device, backend, model_path = load_model(
model_path, resolved_device, defaults["compute_type"], defaults["model"], resolved_device, defaults["compute_type"],
on_status=lambda msg: console.print(msg), strict_device=strict, on_status=lambda msg: console.print(msg), strict_device=strict,
compute_type_explicit=compute_type_explicit,
)
actual_ct = getattr(backend, "actual_compute_type", defaults["compute_type"]) or defaults["compute_type"]
console.print(
f"Модель: [bold]{defaults['model']}[/bold] "
f"Устройство: [bold]{actual_device}[/bold] "
f"Compute: [bold]{actual_ct}[/bold]"
) )
with Status("Подготавливаю запуск...", console=console) as status: with Status("Подготавливаю запуск...", console=console) as status:
tfr = _transcribe_file( tfr = _transcribe_file(
model=model_obj, model=model_obj,
actual_device=actual_device, actual_device=actual_device,
backend=backend,
model_path=model_path,
file_path=validated_file, file_path=validated_file,
model_name=model_path, model_name=defaults["model"],
compute_type=defaults["compute_type"], compute_type=defaults["compute_type"],
language=defaults["language"] if defaults["language"] != "auto" else None, language=defaults["language"] if defaults["language"] != "auto" else None,
on_segment=on_segment if verbose else None, on_segment=on_segment if verbose else None,
@@ -176,12 +187,7 @@ def _run_single(
f"Речь не обнаружена в файле {validated_file.name}", style="yellow" f"Речь не обнаружена в файле {validated_file.name}", style="yellow"
) )
if result.device_used == "cuda": device_info = _format_device_info(result.device_used)
gpu_name = get_gpu_name()
device_info = f"CUDA ({gpu_name or 'Unknown GPU'})"
else:
device_info = "CPU"
language_mode = "detected" if defaults["language"] == "auto" else "forced" language_mode = "detected" if defaults["language"] == "auto" else "forced"
content = format_transcript( content = format_transcript(
@@ -194,7 +200,7 @@ def _run_single(
write_transcript(content, output_path) write_transcript(content, output_path)
elapsed = time.monotonic() - start elapsed = time.monotonic() - start
console.print(f"Транскрипт сохранён: [bold]{output_path}[/bold]", style="green") console.print(f"Транскрипт сохранён: \"{output_path}\"", style="green")
console.print(f" Сегментов: {len(result.segments)} Время: {elapsed:.1f}с") console.print(f" Сегментов: {len(result.segments)} Время: {elapsed:.1f}с")
@@ -203,6 +209,7 @@ def _run_batch(
defaults: dict[str, str], defaults: dict[str, str],
verbose: bool, verbose: bool,
force: bool, force: bool,
compute_type_explicit: bool = False,
) -> None: ) -> None:
"""Трёхфазный батч-пайплайн: prescan → загрузка модели → транскрипция.""" """Трёхфазный батч-пайплайн: prescan → загрузка модели → транскрипция."""
# Phase 1: Prescan — fail-fast + skip до загрузки модели (экономим ~2-5 сек) # Phase 1: Prescan — fail-fast + skip до загрузки модели (экономим ~2-5 сек)
@@ -232,16 +239,14 @@ def _run_batch(
raise SystemExit(1) raise SystemExit(1)
return return
# Phase 2: Load model # Phase 2: Load model (ensure + create в одном вызове)
requested_device = defaults["device"] requested_device = defaults["device"]
resolved_device = detect_device(requested_device) resolved_device = detect_device(requested_device)
strict = requested_device != "auto" strict = requested_device != "auto"
model_path = ensure_model_available( model_obj, actual_device, backend, model_path = load_model(
defaults["model"], on_status=lambda msg: console.print(msg) defaults["model"], resolved_device, defaults["compute_type"],
)
model_obj, actual_device = load_model(
model_path, resolved_device, defaults["compute_type"],
on_status=lambda msg: console.print(msg), strict_device=strict, on_status=lambda msg: console.print(msg), strict_device=strict,
compute_type_explicit=compute_type_explicit,
) )
if actual_device != resolved_device: if actual_device != resolved_device:
@@ -277,8 +282,10 @@ def _run_batch(
tfr = _transcribe_file( tfr = _transcribe_file(
model=model_obj, model=model_obj,
actual_device=actual_device, actual_device=actual_device,
backend=backend,
model_path=model_path,
file_path=file, file_path=file,
model_name=model_path, model_name=defaults["model"],
compute_type=defaults["compute_type"], compute_type=defaults["compute_type"],
language=defaults["language"] if defaults["language"] != "auto" else None, language=defaults["language"] if defaults["language"] != "auto" else None,
on_segment=on_segment if verbose else None, on_segment=on_segment if verbose else None,
@@ -291,8 +298,11 @@ def _run_batch(
f" {file.name}: fallback на {tfr.actual_device} при транскрипции", f" {file.name}: fallback на {tfr.actual_device} при транскрипции",
style="yellow", style="yellow",
) )
# Обновляем после возможного mid-stream fallback на CPU # Обновляем после возможного mid-stream fallback
model_obj, actual_device = tfr.model, tfr.actual_device model_obj = tfr.model
actual_device = tfr.actual_device
backend = tfr.backend
model_path = tfr.model_path
result = tfr.result result = tfr.result
@@ -301,11 +311,7 @@ def _run_batch(
f" Речь не обнаружена: {file.name}", style="yellow" f" Речь не обнаружена: {file.name}", style="yellow"
) )
if result.device_used == "cuda": device_info = _format_device_info(result.device_used)
gpu_name = get_gpu_name()
device_info = f"CUDA ({gpu_name or 'Unknown GPU'})"
else:
device_info = "CPU"
content = format_transcript( content = format_transcript(
result=result, result=result,
+2 -1
View File
@@ -19,11 +19,12 @@ HARDCODED_DEFAULTS: dict[str, str] = {
DEVICE_DEFAULTS: dict[str, dict[str, str]] = { DEVICE_DEFAULTS: dict[str, dict[str, str]] = {
"cuda": {"model": "medium", "compute_type": "float16"}, "cuda": {"model": "medium", "compute_type": "float16"},
"cpu": {"model": "medium", "compute_type": "float32"}, "cpu": {"model": "medium", "compute_type": "float32"},
"openvino": {"model": "medium", "compute_type": "int8"},
} }
# Одно место правды для допустимых ключей конфига # Одно место правды для допустимых ключей конфига
_VALID_KEYS = set(HARDCODED_DEFAULTS) _VALID_KEYS = set(HARDCODED_DEFAULTS)
_VALID_DEVICES = {"auto", "cpu", "cuda"} _VALID_DEVICES = {"auto", "cpu", "cuda", "openvino"}
def find_config_file() -> Path | None: def find_config_file() -> Path | None:
+1 -1
View File
@@ -4,7 +4,7 @@ from dataclasses import dataclass
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
from .transcriber import Segment, TranscribeResult from .types import Segment, TranscribeResult
_PAUSE_THRESHOLD_S = 2.0 # пауза между сегментами для разбиения на абзацы _PAUSE_THRESHOLD_S = 2.0 # пауза между сегментами для разбиения на абзацы
_MAX_PARAGRAPH_S = 60.0 # максимальная длительность абзаца _MAX_PARAGRAPH_S = 60.0 # максимальная длительность абзаца
+81 -188
View File
@@ -1,65 +1,18 @@
"""Обёртка над faster-whisper: загрузка моделей, транскрипция, CUDA fallback.""" """Оркестрация транскрипции: выбор бэкенда, загрузка модели, fallback."""
import warnings import warnings
from collections.abc import Callable from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import Any
# Должен быть ДО импорта faster_whisper / ctranslate2 from local_transcriber.backends import get_backend
from local_transcriber._cuda_bootstrap import ensure_cublas_loadable
ensure_cublas_loadable() # Re-export из types.py для обратной совместимости
from local_transcriber.types import ( # noqa: F401
from faster_whisper import WhisperModel # noqa: E402 Segment,
from huggingface_hub import snapshot_download TranscribeFileResult,
from huggingface_hub.errors import LocalEntryNotFoundError TranscribeResult,
)
MODEL_REPOS = {
"tiny": "Systran/faster-whisper-tiny",
"base": "Systran/faster-whisper-base",
"small": "Systran/faster-whisper-small",
"medium": "Systran/faster-whisper-medium",
"large-v3": "Systran/faster-whisper-large-v3",
}
# allow — фильтр для snapshot_download (какие файлы скачивать из репозитория);
# required — для валидации (что обязано быть после скачивания/в локальной модели)
MODEL_ALLOW_PATTERNS = [
"config.json",
"preprocessor_config.json",
"model.bin",
"tokenizer.json",
"vocabulary.*",
]
MODEL_REQUIRED_FILES = [
"config.json",
"model.bin",
"tokenizer.json",
]
@dataclass
class Segment:
start: float # seconds
end: float # seconds
text: str
@dataclass
class TranscribeResult:
segments: list[Segment]
language: str
language_probability: float
duration: float # seconds
device_used: str # "cpu" / "cuda"
@dataclass
class TranscribeFileResult:
result: TranscribeResult
model: WhisperModel
actual_device: str
def load_model( def load_model(
@@ -68,15 +21,23 @@ def load_model(
compute_type: str, compute_type: str,
on_status: Callable[[str], None] | None = None, on_status: Callable[[str], None] | None = None,
strict_device: bool = False, strict_device: bool = False,
) -> tuple[WhisperModel, str]: compute_type_explicit: bool = False,
"""Загружает модель с CUDA-фолбеком. Возвращает (model, actual_device).""" ) -> tuple[Any, str, Any, str]:
"""Загружает модель: ensure + create с fallback.
Возвращает (model, actual_device, backend, model_path).
compute_type_explicit: True если пользователь явно указал --compute-type.
"""
backend = get_backend(device, compute_type_explicit=compute_type_explicit)
actual_device = device actual_device = device
model_path = backend.ensure_model_available(model_name, compute_type, on_status)
try: try:
_notify_status(on_status, f"Инициализирую модель на {device}...") _notify_status(on_status, f"Инициализирую модель на {device}...")
model = _create_model(model_name, device, compute_type) model = backend.create_model(model_path, device, compute_type)
except (RuntimeError, ValueError) as exc: except (RuntimeError, ValueError) as exc:
# strict — пользователь явно указал устройство, fallback запрещён if device != "cpu" and _is_backend_error(exc, device):
if device != "cpu" and _is_cuda_error(exc):
if strict_device: if strict_device:
raise raise
warnings.warn( warnings.warn(
@@ -85,16 +46,21 @@ def load_model(
stacklevel=2, stacklevel=2,
) )
actual_device = "cpu" actual_device = "cpu"
backend = get_backend("cpu")
model_path = backend.ensure_model_available(model_name, compute_type, on_status)
_notify_status(on_status, "Инициализирую модель на cpu...") _notify_status(on_status, "Инициализирую модель на cpu...")
model = _create_model(model_name, "cpu", compute_type) model = backend.create_model(model_path, "cpu", compute_type)
else: else:
raise raise
return model, actual_device
return model, actual_device, backend, model_path
def _transcribe_file( def _transcribe_file(
model: WhisperModel, model: Any,
actual_device: str, actual_device: str,
backend: Any,
model_path: str,
file_path: Path, file_path: Path,
model_name: str, model_name: str,
compute_type: str, compute_type: str,
@@ -103,39 +69,40 @@ def _transcribe_file(
on_status: Callable[[str], None] | None = None, on_status: Callable[[str], None] | None = None,
strict_device: bool = False, strict_device: bool = False,
) -> TranscribeFileResult: ) -> TranscribeFileResult:
"""Транскрибирует один файл. При mid-stream CUDA fallback перезагружает модель.""" """Транскрибирует один файл. При mid-stream fallback перезагружает модель."""
lang_arg = language if language and language != "auto" else None lang_arg = language if language and language != "auto" else None
try: try:
_notify_status(on_status, "Транскрибирую...") _notify_status(on_status, "Транскрибирую...")
segments, info = _run_transcription(model, file_path, lang_arg, on_segment, on_status) result = backend.transcribe(model, file_path, lang_arg, on_segment, on_status)
result.device_used = actual_device
except (RuntimeError, ValueError) as exc: except (RuntimeError, ValueError) as exc:
# Mid-stream fallback: GPU может упасть с OOM уже во время транскрипции, if actual_device != "cpu" and _is_backend_error(exc, actual_device):
# поэтому перезагружаем модель на CPU и начинаем сначала
if actual_device != "cpu" and _is_cuda_error(exc):
if strict_device: if strict_device:
raise raise
warnings.warn( warnings.warn(
f"CUDA ошибка при транскрипции: {exc}. " f"Ошибка при транскрипции на {actual_device}: {exc}. "
"Переключение на CPU и повтор.", "Переключение на CPU и повтор.",
stacklevel=2, stacklevel=2,
) )
actual_device = "cpu" actual_device = "cpu"
backend = get_backend("cpu")
model_path = backend.ensure_model_available(model_name, compute_type, on_status)
_notify_status(on_status, "Инициализирую модель на cpu...") _notify_status(on_status, "Инициализирую модель на cpu...")
model = _create_model(model_name, "cpu", compute_type) model = backend.create_model(model_path, "cpu", compute_type)
_notify_status(on_status, "Транскрибирую...") _notify_status(on_status, "Транскрибирую...")
segments, info = _run_transcription(model, file_path, lang_arg, on_segment, on_status) result = backend.transcribe(model, file_path, lang_arg, on_segment, on_status)
result.device_used = actual_device
else: else:
raise raise
result = TranscribeResult( return TranscribeFileResult(
segments=segments, result=result,
language=info.language, model=model,
language_probability=info.language_probability, actual_device=actual_device,
duration=info.duration, backend=backend,
device_used=actual_device, model_path=model_path,
) )
return TranscribeFileResult(result=result, model=model, actual_device=actual_device)
def transcribe( def transcribe(
@@ -149,9 +116,13 @@ def transcribe(
strict_device: bool = False, strict_device: bool = False,
) -> TranscribeResult: ) -> TranscribeResult:
"""High-level API: загрузка модели + транскрипция за один вызов.""" """High-level API: загрузка модели + транскрипция за один вызов."""
model, actual_device = load_model(model_name, device, compute_type, on_status, strict_device) model, actual_device, backend, model_path = load_model(
model_name, device, compute_type, on_status, strict_device,
compute_type_explicit=True, # Python API — caller explicitly chose compute_type
)
tfr = _transcribe_file( tfr = _transcribe_file(
model, actual_device, file_path, model_name, compute_type, model, actual_device, backend, model_path,
file_path, model_name, compute_type,
language, on_segment, on_status, strict_device, language, on_segment, on_status, strict_device,
) )
return tfr.result return tfr.result
@@ -159,127 +130,49 @@ def transcribe(
def ensure_model_available( def ensure_model_available(
model_name: str, model_name: str,
device: str = "cpu",
compute_type: str | None = None,
on_status: Callable[[str], None] | None = None, on_status: Callable[[str], None] | None = None,
) -> str: ) -> str:
"""Резолвит alias модели в repo_id и гарантирует наличие файлов. """Публичный helper: гарантирует наличие модели для указанного бэкенда."""
from local_transcriber.config import DEVICE_DEFAULTS, HARDCODED_DEFAULTS
Стратегия: cache-first (``local_files_only=True``), затем download. if compute_type is None:
Два вызова ``snapshot_download`` — чтобы не лезть в сеть, если модель уже в кэше. device_defs = DEVICE_DEFAULTS.get(device, {})
""" compute_type = device_defs.get("compute_type", HARDCODED_DEFAULTS["compute_type"])
local_path = Path(model_name).expanduser() explicit = False
if local_path.is_dir(): else:
_validate_model_dir(local_path) explicit = True
return str(local_path) backend = get_backend(device, compute_type_explicit=explicit)
return backend.ensure_model_available(model_name, compute_type, on_status)
repo_id = _resolve_model_repo(model_name)
try:
_notify_status(on_status, f"Проверяю кэш модели {model_name}...")
cached_path = Path(_snapshot_download(repo_id, local_files_only=True))
_validate_model_dir(cached_path)
return str(cached_path)
except LocalEntryNotFoundError:
pass
except ValueError:
_notify_status(on_status, f"Кэш модели {model_name} неполный, докачиваю...")
_notify_status(on_status, f"Скачиваю модель {model_name} из Hugging Face...")
downloaded_path = Path(_snapshot_download(repo_id, local_files_only=False))
_validate_model_dir(downloaded_path)
return str(downloaded_path)
def _run_transcription(model, file_path, lang_arg, on_segment, on_status=None):
"""Run model.transcribe and iterate segments. Returns (segments, info)."""
segment_generator, info = model.transcribe(str(file_path), language=lang_arg)
total_duration = info.duration
segments: list[Segment] = []
for raw_seg in segment_generator:
seg = Segment(start=raw_seg.start, end=raw_seg.end, text=raw_seg.text)
if on_segment is not None:
on_segment(seg)
segments.append(seg)
_notify_status(
on_status,
f"Транскрибирую... {_fmt_time(seg.end)} / {_fmt_time(total_duration)}"
f" [{len(segments)} сегм.]",
)
return segments, info
def _create_model(model_name: str, device: str, compute_type: str):
try:
return WhisperModel(model_name, device=device, compute_type=compute_type)
except ImportError as exc:
# WhisperModel при инициализации может загружать файлы через HF Hub;
# если в системе настроен SOCKS proxy, но socksio не установлен,
# HF Hub бросает ImportError — оборачиваем в понятное сообщение
if _is_missing_socksio_error(exc):
raise RuntimeError(
"Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, "
"нужная для загрузки модели из Hugging Face через proxy. "
"Обновите окружение: `uv sync`."
) from exc
raise
def _is_cuda_error(exc: BaseException) -> bool: def _is_cuda_error(exc: BaseException) -> bool:
"""Проверка CUDA ошибок — используется в cli.py для Windows-диагностики."""
msg = str(exc).lower() msg = str(exc).lower()
return any(k in msg for k in ("cuda", "cublas", "cudnn", "out of memory")) return any(k in msg for k in ("cuda", "cublas", "cudnn", "out of memory"))
def _is_missing_socksio_error(exc: BaseException) -> bool: def _is_backend_error(exc: BaseException, device: str) -> bool:
msg = str(exc).lower() """Определяет, связана ли ошибка с конкретным бэкендом (а не с пользовательскими данными)."""
return "socks proxy" in msg and "socksio" in msg if device in ("cuda", "cpu"):
return _is_cuda_error(exc)
if device == "openvino":
return _is_openvino_error(exc)
return False
def _fmt_time(seconds: float) -> str: def _is_openvino_error(exc: BaseException) -> bool:
m, s = divmod(int(seconds), 60) """Проверка ошибок OpenVINO runtime.
h, m = divmod(m, 60)
return f"{h}:{m:02d}:{s:02d}" if h else f"{m:02d}:{s:02d}" OpenVINO runtime кидает RuntimeError с разнообразными сообщениями
(openvino, ov_, inference, plugins, src/...). Пользовательские ошибки
(файл не найден, неверный формат) приходят как FileNotFoundError/ValueError
и не попадают сюда. Поэтому для RuntimeError считаем это backend failure.
"""
return isinstance(exc, RuntimeError)
def _notify_status(on_status: Callable[[str], None] | None, message: str) -> None: def _notify_status(on_status: Callable[[str], None] | None, message: str) -> None:
if on_status is not None: if on_status is not None:
on_status(message) on_status(message)
def _resolve_model_repo(model_name: str) -> str:
if "/" in model_name:
return model_name
repo_id = MODEL_REPOS.get(model_name)
if repo_id is None:
expected = ", ".join(MODEL_REPOS)
raise ValueError(f"Неподдерживаемая модель '{model_name}'. Ожидалось одно из: {expected}")
return repo_id
def _snapshot_download(repo_id: str, local_files_only: bool) -> str:
try:
return snapshot_download(
repo_id,
local_files_only=local_files_only,
allow_patterns=MODEL_ALLOW_PATTERNS,
)
except ImportError as exc:
if _is_missing_socksio_error(exc):
raise RuntimeError(
"Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, "
"нужная для загрузки модели из Hugging Face через proxy. "
"Обновите окружение: `uv sync`."
) from exc
raise
def _validate_model_dir(model_dir: Path) -> None:
missing = [
filename for filename in MODEL_REQUIRED_FILES if not (model_dir / filename).exists()
]
if not any(model_dir.glob("vocabulary.*")):
missing.append("vocabulary.*")
if missing:
missing_str = ", ".join(missing)
raise ValueError(f"Неполная локальная модель в '{model_dir}': отсутствуют {missing_str}")
+34
View File
@@ -0,0 +1,34 @@
"""Общие типы данных для всех бэкендов транскрипции."""
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
@dataclass
class Segment:
start: float # seconds
end: float # seconds
text: str
@dataclass
class TranscribeResult:
segments: list[Segment]
language: str
language_probability: float
duration: float # seconds
device_used: str # "cpu" / "cuda" / "openvino"
@dataclass
class TranscribeFileResult:
result: TranscribeResult
model: Any # backend-specific model handle
actual_device: str
backend: Any = None # backend instance (для переиспользования в батче)
model_path: str = "" # путь к модели (меняется при cross-backend fallback)
StatusCallback = Callable[[str], None] | None
SegmentCallback = Callable[[Segment], None] | None
+17 -2
View File
@@ -1,6 +1,7 @@
"""Утилиты для валидации входных файлов, определения устройства и работы с путями.""" """Утилиты для валидации входных файлов, определения устройства и работы с путями."""
import glob import glob
import platform
import shutil import shutil
import subprocess import subprocess
import warnings import warnings
@@ -15,16 +16,30 @@ SUPPORTED_EXTENSIONS = {
def detect_device(requested: str = "auto") -> str: def detect_device(requested: str = "auto") -> str:
"""Определяет устройство для вычислений. """Определяет устройство для вычислений.
При ``requested="auto"`` проверяет наличие ``nvidia-smi`` в PATH При ``requested="auto"`` проверяет: CUDA → OpenVINO → CPU.
и возвращает ``"cuda"`` или ``"cpu"``. Явное значение возвращается как есть. Явное значение возвращается как есть.
""" """
if requested != "auto": if requested != "auto":
return requested return requested
if shutil.which("nvidia-smi") is not None: if shutil.which("nvidia-smi") is not None:
return "cuda" return "cuda"
if _is_openvino_available():
return "openvino"
return "cpu" return "cpu"
def _is_openvino_available() -> bool:
"""Проверяет доступность OpenVINO: x86/AMD64 архитектура + пакет установлен."""
if platform.machine().lower() not in {"x86_64", "amd64"}:
return False
try:
import openvino_genai # noqa: F401
return True
except ImportError:
return False
def get_gpu_name() -> str | None: def get_gpu_name() -> str | None:
"""Возвращает название GPU через ``nvidia-smi`` (для метаданных транскрипта).""" """Возвращает название GPU через ``nvidia-smi`` (для метаданных транскрипта)."""
try: try:
+233
View File
@@ -0,0 +1,233 @@
"""Тесты для OpenVINO бэкенда."""
from pathlib import Path
from unittest.mock import MagicMock, patch
import numpy as np
import pytest
from local_transcriber.backends.openvino import (
MODEL_REPOS,
OpenVINOBackend,
_validate_model_dir,
)
from local_transcriber.types import Segment
# === _resolve_repo ===
def test_resolve_repo_exact_match():
backend = OpenVINOBackend(compute_type_explicit=True)
assert backend._resolve_repo("medium", "int8") == ("OpenVINO/whisper-medium-int8-ov", "int8")
def test_resolve_repo_large_v3_fp16():
backend = OpenVINOBackend(compute_type_explicit=True)
assert backend._resolve_repo("large-v3", "fp16") == ("OpenVINO/whisper-large-v3-fp16-ov", "fp16")
def test_resolve_repo_explicit_unsupported_pair_raises():
"""Явный --compute-type с несуществующей парой → ошибка."""
backend = OpenVINOBackend(compute_type_explicit=True)
with pytest.raises(ValueError, match="недоступна с compute_type='fp16'"):
backend._resolve_repo("medium", "fp16")
def test_resolve_repo_explicit_unknown_model_raises():
backend = OpenVINOBackend(compute_type_explicit=True)
with pytest.raises(ValueError, match="не найдена для OpenVINO"):
backend._resolve_repo("distil-large-v3", "int8")
def test_resolve_repo_implicit_fallback():
"""Неявный compute_type: если int8 недоступен для base, fallback на fp16."""
backend = OpenVINOBackend(compute_type_explicit=False)
# base + int8 не существует, но base + fp16 есть
assert backend._resolve_repo("base", "int8") == ("OpenVINO/whisper-base-fp16-ov", "fp16")
def test_resolve_repo_implicit_large_v3_prefers_fp16():
"""Неявный compute_type: large-v3 автоматически получает fp16."""
backend = OpenVINOBackend(compute_type_explicit=False)
# Дефолт int8, но для large-v3 override на fp16
assert backend._resolve_repo("large-v3", "int8") == ("OpenVINO/whisper-large-v3-fp16-ov", "fp16")
def test_resolve_repo_explicit_large_v3_int8_respected():
"""Явный --compute-type int8 для large-v3 → уважается."""
backend = OpenVINOBackend(compute_type_explicit=True)
assert backend._resolve_repo("large-v3", "int8") == ("OpenVINO/whisper-large-v3-int8-ov", "int8")
# === ensure_model_available ===
@patch("local_transcriber.backends.openvino.snapshot_download")
def test_ensure_model_available_cache_hit(mock_download, tmp_path):
model_dir = tmp_path / "model"
model_dir.mkdir()
(model_dir / "openvino_encoder_model.xml").write_text("<xml/>")
(model_dir / "openvino_decoder_model.xml").write_text("<xml/>")
mock_download.return_value = str(model_dir)
backend = OpenVINOBackend(compute_type_explicit=True)
result = backend.ensure_model_available("medium", "int8")
assert result == str(model_dir)
mock_download.assert_called_once()
assert mock_download.call_args.kwargs["local_files_only"] is True
@patch("local_transcriber.backends.openvino.snapshot_download")
def test_ensure_model_available_downloads(mock_download, tmp_path):
from huggingface_hub.errors import LocalEntryNotFoundError
model_dir = tmp_path / "downloaded"
model_dir.mkdir()
(model_dir / "openvino_encoder_model.xml").write_text("<xml/>")
(model_dir / "openvino_decoder_model.xml").write_text("<xml/>")
mock_download.side_effect = [
LocalEntryNotFoundError("not cached"),
str(model_dir),
]
backend = OpenVINOBackend(compute_type_explicit=True)
statuses: list[str] = []
result = backend.ensure_model_available("medium", "int8", on_status=statuses.append)
assert result == str(model_dir)
assert any("Скачиваю" in s for s in statuses)
# === create_model ===
def test_create_model():
mock_ov = MagicMock()
mock_pipeline = MagicMock()
mock_ov.WhisperPipeline.return_value = mock_pipeline
backend = OpenVINOBackend()
with patch.dict("sys.modules", {"openvino_genai": mock_ov}):
model = backend.create_model("/path/to/model", "openvino", "int8")
mock_ov.WhisperPipeline.assert_called_once_with("/path/to/model", "CPU")
assert model is mock_pipeline
# === transcribe ===
def test_transcribe_maps_chunks_to_segments():
"""Проверяет маппинг chunks → Segment[] и формат языка."""
backend = OpenVINOBackend()
mock_model = MagicMock()
chunk1 = MagicMock()
chunk1.start_ts = 0.0
chunk1.end_ts = 3.5
chunk1.text = " Привет мир"
chunk2 = MagicMock()
chunk2.start_ts = 3.5
chunk2.end_ts = 7.0
chunk2.text = " Тестовый сегмент"
mock_result = MagicMock()
mock_result.chunks = [chunk1, chunk2]
mock_model.generate.return_value = mock_result
raw_audio = np.zeros(16000 * 10, dtype=np.float32) # 10 секунд
with patch("faster_whisper.decode_audio", return_value=raw_audio):
result = backend.transcribe(
mock_model, Path("test.mp3"), language="ru",
)
assert len(result.segments) == 2
assert result.segments[0].text == " Привет мир"
assert result.segments[0].start == 0.0
assert result.segments[0].end == 3.5
assert result.duration == 10.0
# Проверяем формат языка для OpenVINO GenAI
call_kwargs = mock_model.generate.call_args
assert call_kwargs.kwargs["language"] == "<|ru|>"
assert call_kwargs.kwargs["return_timestamps"] is True
def test_transcribe_calls_tolist():
"""raw_speech передаётся как list, не ndarray."""
backend = OpenVINOBackend()
mock_model = MagicMock()
mock_result = MagicMock()
mock_result.chunks = []
mock_model.generate.return_value = mock_result
raw_audio = np.zeros(160, dtype=np.float32)
with patch("faster_whisper.decode_audio", return_value=raw_audio):
backend.transcribe(mock_model, Path("test.mp3"), language=None)
call_args = mock_model.generate.call_args[0][0]
assert isinstance(call_args, list)
def test_transcribe_no_language_auto():
"""Без указания языка — не передаём language в generate."""
backend = OpenVINOBackend()
mock_model = MagicMock()
mock_result = MagicMock()
mock_result.chunks = []
mock_model.generate.return_value = mock_result
raw_audio = np.zeros(160, dtype=np.float32)
with patch("faster_whisper.decode_audio", return_value=raw_audio):
result = backend.transcribe(mock_model, Path("test.mp3"), language=None)
call_kwargs = mock_model.generate.call_args.kwargs
assert "language" not in call_kwargs
assert result.language == "auto"
assert result.language_probability == 0.0
def test_transcribe_calls_on_segment():
backend = OpenVINOBackend()
mock_model = MagicMock()
chunk = MagicMock()
chunk.start_ts = 0.0
chunk.end_ts = 2.0
chunk.text = " Test"
mock_result = MagicMock()
mock_result.chunks = [chunk]
mock_model.generate.return_value = mock_result
raw_audio = np.zeros(16000, dtype=np.float32)
callback = MagicMock()
with patch("faster_whisper.decode_audio", return_value=raw_audio):
backend.transcribe(
mock_model, Path("test.mp3"), language="en", on_segment=callback,
)
callback.assert_called_once()
seg = callback.call_args[0][0]
assert isinstance(seg, Segment)
assert seg.text == " Test"
# === _validate_model_dir ===
def test_validate_model_dir_ok(tmp_path):
(tmp_path / "openvino_encoder_model.xml").write_text("<xml/>")
(tmp_path / "openvino_decoder_model.xml").write_text("<xml/>")
_validate_model_dir(tmp_path) # should not raise
def test_validate_model_dir_missing(tmp_path):
(tmp_path / "openvino_encoder_model.xml").write_text("<xml/>")
with pytest.raises(ValueError, match="openvino_decoder_model.xml"):
_validate_model_dir(tmp_path)
+103 -94
View File
@@ -24,12 +24,21 @@ def _make_model():
return MagicMock(name="WhisperModel") return MagicMock(name="WhisperModel")
def _make_tfr(result=None, model=None, actual_device="cpu"): def _make_backend():
return MagicMock(name="Backend")
def _make_tfr(result=None, model=None, actual_device="cpu", backend=None, model_path="/models/medium"):
if result is None: if result is None:
result = _make_result() result = _make_result()
if model is None: if model is None:
model = _make_model() model = _make_model()
return TranscribeFileResult(result=result, model=model, actual_device=actual_device) if backend is None:
backend = _make_backend()
return TranscribeFileResult(
result=result, model=model, actual_device=actual_device,
backend=backend, model_path=model_path,
)
def _single_patches(result=None, tmp_file=None, actual_device="cpu"): def _single_patches(result=None, tmp_file=None, actual_device="cpu"):
@@ -37,13 +46,13 @@ def _single_patches(result=None, tmp_file=None, actual_device="cpu"):
if result is None: if result is None:
result = _make_result(device_used=actual_device) result = _make_result(device_used=actual_device)
model = _make_model() model = _make_model()
tfr = TranscribeFileResult(result=result, model=model, actual_device=actual_device) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device=actual_device, backend=backend)
return [ return [
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=tmp_file), patch("local_transcriber.cli.validate_input_file", return_value=tmp_file),
patch("local_transcriber.cli.detect_device", return_value=actual_device), patch("local_transcriber.cli.detect_device", return_value=actual_device),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, actual_device, backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, actual_device)),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
] ]
@@ -54,7 +63,7 @@ def test_cli_happy_path_exit_code_zero(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
patches = _single_patches(tmp_file=audio) patches = _single_patches(tmp_file=audio)
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6]: with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5]:
out = runner.invoke(app, [str(audio)]) out = runner.invoke(app, [str(audio)])
assert out.exit_code == 0 assert out.exit_code == 0
@@ -65,22 +74,22 @@ def test_cli_default_options_passed_to_transcribe(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
runner.invoke(app, [str(audio)]) runner.invoke(app, [str(audio)])
call_kwargs = mock_transcribe_file.call_args[1] call_kwargs = mock_transcribe_file.call_args[1]
assert call_kwargs["model_name"] == "/models/medium" assert call_kwargs["model_name"] == "medium"
assert call_kwargs["compute_type"] == "float32" assert call_kwargs["compute_type"] == "float32"
assert call_kwargs["language"] == "ru" assert call_kwargs["language"] == "ru"
assert call_kwargs["on_segment"] is None # verbose=False assert call_kwargs["on_segment"] is None # verbose=False
@@ -91,15 +100,15 @@ def test_cli_custom_options(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result(device_used="cuda") result = _make_result(device_used="cuda")
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model, actual_device="cuda") backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device="cuda", backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/small"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/small")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"), patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"),
@@ -113,7 +122,7 @@ def test_cli_custom_options(tmp_path):
]) ])
call_kwargs = mock_transcribe_file.call_args[1] call_kwargs = mock_transcribe_file.call_args[1]
assert call_kwargs["model_name"] == "/models/small" assert call_kwargs["model_name"] == "small"
assert call_kwargs["language"] == "ru" assert call_kwargs["language"] == "ru"
assert call_kwargs["compute_type"] == "float16" assert call_kwargs["compute_type"] == "float16"
@@ -123,15 +132,15 @@ def test_cli_verbose_passes_on_segment_callback(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -148,7 +157,7 @@ def test_cli_empty_speech_warning(tmp_path):
result = _make_result(segments=[]) result = _make_result(segments=[])
patches = _single_patches(result=result, tmp_file=audio) patches = _single_patches(result=result, tmp_file=audio)
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6]: with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5]:
out = runner.invoke(app, [str(audio)]) out = runner.invoke(app, [str(audio)])
assert out.exit_code == 0 assert out.exit_code == 0
@@ -162,14 +171,14 @@ def test_cli_default_output_path(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript", mock_write), patch("local_transcriber.cli.write_transcript", mock_write),
): ):
@@ -187,14 +196,14 @@ def test_cli_custom_output_path(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript", mock_write), patch("local_transcriber.cli.write_transcript", mock_write),
): ):
@@ -209,15 +218,15 @@ def test_cli_passes_status_callback_to_transcribe(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -228,28 +237,27 @@ def test_cli_passes_status_callback_to_transcribe(tmp_path):
assert callable(call_kwargs["on_status"]) assert callable(call_kwargs["on_status"])
def test_cli_resolves_model_before_transcribe(tmp_path): def test_cli_load_model_called_with_model_name(tmp_path):
"""load_model receives model name from defaults, handles ensure internally."""
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
mock_transcribe_file = MagicMock(return_value=tfr) tfr = _make_tfr(result=result, model=model, backend=backend)
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/large-v3"))
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3") as mock_ensure, patch("local_transcriber.cli.load_model", mock_load_model),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
runner.invoke(app, [str(audio), "--model", "large-v3"]) runner.invoke(app, [str(audio), "--model", "large-v3"])
mock_ensure.assert_called_once() assert mock_load_model.call_args[0][0] == "large-v3"
call_kwargs = mock_transcribe_file.call_args[1]
assert call_kwargs["model_name"] == "/models/large-v3"
def test_cli_windows_cuda_diagnostic(tmp_path): def test_cli_windows_cuda_diagnostic(tmp_path):
@@ -257,13 +265,13 @@ def test_cli_windows_cuda_diagnostic(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")),
patch("local_transcriber.cli.sys") as mock_sys, patch("local_transcriber.cli.sys") as mock_sys,
): ):
@@ -280,13 +288,13 @@ def test_cli_linux_cuda_error_no_windows_hint(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")),
patch("local_transcriber.cli.sys") as mock_sys, patch("local_transcriber.cli.sys") as mock_sys,
): ):
@@ -303,14 +311,14 @@ def test_cli_device_fallback_warning(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result(device_used="cpu") result = _make_result(device_used="cpu")
model = _make_model() model = _make_model()
tfr = TranscribeFileResult(result=result, model=model, actual_device="cpu") backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device="cpu", backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -325,15 +333,15 @@ def test_cli_strict_device_passed_to_transcribe(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result(device_used="cuda") result = _make_result(device_used="cuda")
model = _make_model() model = _make_model()
tfr = TranscribeFileResult(result=result, model=model, actual_device="cuda") backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device="cuda", backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"), patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"),
@@ -344,15 +352,14 @@ def test_cli_strict_device_passed_to_transcribe(tmp_path):
mock_transcribe_file.reset_mock() mock_transcribe_file.reset_mock()
result_cpu = _make_result(device_used="cpu") result_cpu = _make_result(device_used="cpu")
tfr_cpu = TranscribeFileResult(result=result_cpu, model=model, actual_device="cpu") tfr_cpu = _make_tfr(result=result_cpu, model=model, backend=backend)
mock_transcribe_file.return_value = tfr_cpu mock_transcribe_file.return_value = tfr_cpu
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -366,13 +373,13 @@ def test_cli_keyboard_interrupt(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=KeyboardInterrupt), patch("local_transcriber.cli._transcribe_file", side_effect=KeyboardInterrupt),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -399,13 +406,13 @@ def test_cli_unexpected_error_verbose_traceback(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -420,13 +427,13 @@ def test_cli_unexpected_error_no_verbose_hint(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -448,14 +455,14 @@ def test_cli_batch_two_files(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -470,19 +477,18 @@ def test_cli_batch_skips_existing(tmp_path):
b = tmp_path / "b.mp3" b = tmp_path / "b.mp3"
a.write_bytes(b"fake") a.write_bytes(b"fake")
b.write_bytes(b"fake") b.write_bytes(b"fake")
# Create transcript for a
(tmp_path / "a-transcript.md").write_text("existing") (tmp_path / "a-transcript.md").write_text("existing")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -524,14 +530,14 @@ def test_cli_batch_force_overwrites(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -550,7 +556,8 @@ def test_cli_batch_per_file_error(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
call_count = 0 call_count = 0
def transcribe_side_effect(**kwargs): def transcribe_side_effect(**kwargs):
@@ -564,8 +571,7 @@ def test_cli_batch_per_file_error(tmp_path):
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=transcribe_side_effect), patch("local_transcriber.cli._transcribe_file", side_effect=transcribe_side_effect),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -584,7 +590,8 @@ def test_cli_batch_invalid_in_prescan(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
def validate_side_effect(p): def validate_side_effect(p):
if not p.exists(): if not p.exists():
@@ -595,8 +602,7 @@ def test_cli_batch_invalid_in_prescan(tmp_path):
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=validate_side_effect), patch("local_transcriber.cli.validate_input_file", side_effect=validate_side_effect),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -634,42 +640,46 @@ def test_cli_config_applied(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
result = _make_result() result = _make_result()
tfr = _make_tfr(result=result, model=model) tfr = _make_tfr(result=result, model=model, backend=backend)
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/tiny"))
with ( with (
patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}), patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/tiny") as mock_ensure, patch("local_transcriber.cli.load_model", mock_load_model),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
runner.invoke(app, [str(audio)]) runner.invoke(app, [str(audio)])
mock_ensure.assert_called_once_with("tiny", on_status=mock_ensure.call_args[1]["on_status"]) # load_model receives model name from config
assert mock_load_model.call_args[0][0] == "tiny"
def test_cli_cli_overrides_config(tmp_path): def test_cli_cli_overrides_config(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
result = _make_result() result = _make_result()
tfr = _make_tfr(result=result, model=model) tfr = _make_tfr(result=result, model=model, backend=backend)
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/small"))
with ( with (
patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}), patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/small") as mock_ensure, patch("local_transcriber.cli.load_model", mock_load_model),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
runner.invoke(app, [str(audio), "--model", "small"]) runner.invoke(app, [str(audio), "--model", "small"])
mock_ensure.assert_called_once_with("small", on_status=mock_ensure.call_args[1]["on_status"]) # CLI --model overrides config
assert mock_load_model.call_args[0][0] == "small"
def test_cli_batch_fallback_warning(tmp_path): def test_cli_batch_fallback_warning(tmp_path):
@@ -681,14 +691,14 @@ def test_cli_batch_fallback_warning(tmp_path):
result = _make_result(device_used="cpu") result = _make_result(device_used="cpu")
model = _make_model() model = _make_model()
tfr = TranscribeFileResult(result=result, model=model, actual_device="cpu") backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device="cpu", backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -707,15 +717,15 @@ def test_cli_batch_empty_speech_warning(tmp_path):
result_empty = _make_result(segments=[]) result_empty = _make_result(segments=[])
result_ok = _make_result() result_ok = _make_result()
model = _make_model() model = _make_model()
tfr_empty = _make_tfr(result=result_empty, model=model) backend = _make_backend()
tfr_ok = _make_tfr(result=result_ok, model=model) tfr_empty = _make_tfr(result=result_empty, model=model, backend=backend)
tfr_ok = _make_tfr(result=result_ok, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_empty, tfr_ok]), patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_empty, tfr_ok]),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -735,17 +745,16 @@ def test_cli_batch_midstream_fallback_warning(tmp_path):
model_gpu = _make_model() model_gpu = _make_model()
model_cpu = _make_model() model_cpu = _make_model()
backend = _make_backend()
result = _make_result(device_used="cpu") result = _make_result(device_used="cpu")
# First file triggers mid-stream fallback tfr_fallback = _make_tfr(result=result, model=model_cpu, actual_device="cpu", backend=backend)
tfr_fallback = TranscribeFileResult(result=result, model=model_cpu, actual_device="cpu") tfr_ok = _make_tfr(result=result, model=model_cpu, actual_device="cpu", backend=backend)
tfr_ok = TranscribeFileResult(result=result, model=model_cpu, actual_device="cpu")
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda")),
patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_fallback, tfr_ok]), patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_fallback, tfr_ok]),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -763,14 +772,14 @@ def test_cli_batch_model_loaded_once(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
mock_load_model = MagicMock(return_value=(model, "cpu")) tfr = _make_tfr(result=result, model=model, backend=backend)
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/medium"))
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
patch("local_transcriber.cli.load_model", mock_load_model), patch("local_transcriber.cli.load_model", mock_load_model),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
+15
View File
@@ -128,3 +128,18 @@ def test_apply_device_defaults_config_overrides():
result = apply_device_defaults(defaults, "cuda", cli, config) result = apply_device_defaults(defaults, "cuda", cli, config)
assert result["model"] == "small" assert result["model"] == "small"
assert result["compute_type"] == "int8" assert result["compute_type"] == "int8"
def test_load_config_openvino_device(tmp_path):
config = tmp_path / "config.toml"
config.write_text('device = "openvino"\n')
result = load_config(config)
assert result == {"device": "openvino"}
def test_apply_device_defaults_openvino():
defaults = {"model": "medium", "language": "ru", "device": "auto", "compute_type": "float32"}
cli = {"model": None, "language": None, "device": None, "compute_type": None}
result = apply_device_defaults(defaults, "openvino", cli, {})
assert result["model"] == "medium"
assert result["compute_type"] == "int8"
+348 -277
View File
@@ -1,9 +1,7 @@
from collections.abc import Generator
from pathlib import Path from pathlib import Path
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import pytest import pytest
from huggingface_hub.errors import LocalEntryNotFoundError
from local_transcriber.transcriber import ( from local_transcriber.transcriber import (
Segment, Segment,
@@ -15,24 +13,53 @@ from local_transcriber.transcriber import (
) )
def _make_raw_segments(count: int) -> list: # === Helpers ===
"""Create mock raw segments as returned by faster-whisper."""
segments = []
for i in range(count):
seg = MagicMock()
seg.start = float(i * 5)
seg.end = float(i * 5 + 4)
seg.text = f" Segment {i}"
segments.append(seg)
return segments
def _make_info(language: str = "ru", probability: float = 0.95, duration: float = 60.0): def _make_result(
info = MagicMock() count: int = 2,
info.language = language language: str = "ru",
info.language_probability = probability probability: float = 0.95,
info.duration = duration duration: float = 60.0,
return info device_used: str = "cpu",
) -> TranscribeResult:
segments = [
Segment(start=float(i * 5), end=float(i * 5 + 4), text=f" Segment {i}")
for i in range(count)
]
return TranscribeResult(
segments=segments,
language=language,
language_probability=probability,
duration=duration,
device_used=device_used,
)
def _make_backend(
model=None,
transcribe_result=None,
create_model_error=None,
transcribe_error=None,
model_path="/mock/model",
):
"""Создаёт mock-бэкенд с настраиваемым поведением."""
backend = MagicMock()
backend.ensure_model_available.return_value = model_path
if create_model_error:
backend.create_model.side_effect = create_model_error
else:
backend.create_model.return_value = model or MagicMock()
if transcribe_error:
backend.transcribe.side_effect = transcribe_error
elif transcribe_result:
backend.transcribe.return_value = transcribe_result
else:
backend.transcribe.return_value = _make_result()
return backend
def _create_model_dir(path: Path) -> Path: def _create_model_dir(path: Path) -> Path:
@@ -45,14 +72,14 @@ def _create_model_dir(path: Path) -> Path:
return path return path
@patch("local_transcriber.transcriber.WhisperModel") # === transcribe() tests ===
def test_transcribe_collects_segments(mock_model_cls):
raw_segments = _make_raw_segments(3)
info = _make_info()
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info) @patch("local_transcriber.transcriber.get_backend")
mock_model_cls.return_value = instance def test_transcribe_collects_segments(mock_get_backend):
result_data = _make_result(count=3)
backend = _make_backend(transcribe_result=result_data)
mock_get_backend.return_value = backend
result = transcribe( result = transcribe(
file_path=Path("test.mp3"), file_path=Path("test.mp3"),
@@ -68,14 +95,11 @@ def test_transcribe_collects_segments(mock_model_cls):
assert result.duration == 60.0 assert result.duration == 60.0
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_calls_on_segment(mock_model_cls): def test_transcribe_calls_on_segment(mock_get_backend):
raw_segments = _make_raw_segments(3) result_data = _make_result(count=3)
info = _make_info() backend = _make_backend(transcribe_result=result_data)
mock_get_backend.return_value = backend
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
mock_model_cls.return_value = instance
callback = MagicMock() callback = MagicMock()
@@ -86,28 +110,24 @@ def test_transcribe_calls_on_segment(mock_model_cls):
on_segment=callback, on_segment=callback,
) )
assert callback.call_count == 3 # on_segment is passed through to backend.transcribe
# Each call should receive a Segment instance call_args = backend.transcribe.call_args
for call_args in callback.call_args_list: assert call_args.kwargs.get("on_segment") is callback or call_args[0][3] is callback
seg = call_args[0][0]
assert isinstance(seg, Segment)
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_cuda_fallback(mock_model_cls): def test_transcribe_cuda_fallback(mock_get_backend):
raw_segments = _make_raw_segments(2) """CUDA error at init -> fallback на CPU."""
info = _make_info() cuda_backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
cpu_backend = _make_backend(
transcribe_result=_make_result(count=2, device_used="cpu"),
model_path="/mock/cpu/model",
)
# First call (cuda) raises, second call (cpu) succeeds def backend_for_device(device, **kwargs):
cpu_instance = MagicMock() return cuda_backend if device == "cuda" else cpu_backend
cpu_instance.transcribe.return_value = (iter(raw_segments), info)
def model_side_effect(model_name, device, compute_type): mock_get_backend.side_effect = backend_for_device
if device == "cuda":
raise RuntimeError("CUDA out of memory")
return cpu_instance
mock_model_cls.side_effect = model_side_effect
with pytest.warns(UserWarning, match="Переключение на CPU"): with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe( result = transcribe(
@@ -120,14 +140,12 @@ def test_transcribe_cuda_fallback(mock_model_cls):
assert len(result.segments) == 2 assert len(result.segments) == 2
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_device_used(mock_model_cls): def test_transcribe_device_used(mock_get_backend):
raw_segments = _make_raw_segments(1) backend = _make_backend(
info = _make_info() transcribe_result=_make_result(count=1, device_used="cuda"),
)
instance = MagicMock() mock_get_backend.return_value = backend
instance.transcribe.return_value = (iter(raw_segments), info)
mock_model_cls.return_value = instance
result = transcribe( result = transcribe(
file_path=Path("test.mp3"), file_path=Path("test.mp3"),
@@ -136,98 +154,47 @@ def test_transcribe_device_used(mock_model_cls):
) )
assert result.device_used == "cuda" assert result.device_used == "cuda"
mock_model_cls.assert_called_once_with("tiny", device="cuda", compute_type="int8") backend.create_model.assert_called_once()
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_cuda_fallback_on_transcribe_call(mock_model_cls): def test_transcribe_cuda_fallback_on_transcribe_call(mock_get_backend):
"""CUDA error in model.transcribe() (not __init__) triggers CPU fallback.""" """CUDA error in transcribe (not init) triggers CPU fallback."""
raw_segments = _make_raw_segments(2) cuda_backend = _make_backend(
info = _make_info() transcribe_error=RuntimeError("CUDA error during transcription"),
)
cuda_instance = MagicMock() cpu_backend = _make_backend(
cuda_instance.transcribe.side_effect = RuntimeError("CUDA error during transcription") transcribe_result=_make_result(count=2, device_used="cpu"),
model_path="/mock/cpu/model",
cpu_instance = MagicMock()
cpu_instance.transcribe.return_value = (iter(raw_segments), info)
call_count = 0
def model_side_effect(model_name, device, compute_type):
nonlocal call_count
call_count += 1
if device == "cuda":
return cuda_instance
return cpu_instance
mock_model_cls.side_effect = model_side_effect
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_midstream_fallback_no_duplicate_callbacks(mock_model_cls):
"""on_segment is not called for partial GPU segments on mid-stream fallback."""
info = _make_info()
# GPU iterator: yields 1 segment then raises CUDA error
def _gpu_generator():
seg = MagicMock()
seg.start = 0.0
seg.end = 4.0
seg.text = " GPU seg"
yield seg
raise RuntimeError("CUDA out of memory mid-stream")
cuda_instance = MagicMock()
cuda_instance.transcribe.return_value = (_gpu_generator(), info)
cpu_segments = _make_raw_segments(2)
cpu_instance = MagicMock()
cpu_instance.transcribe.return_value = (iter(cpu_segments), info)
def model_side_effect(model_name, device, compute_type):
if device == "cuda":
return cuda_instance
return cpu_instance
mock_model_cls.side_effect = model_side_effect
callback = MagicMock()
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
on_segment=callback,
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
# callback: 1 from partial GPU pass + 2 from full CPU pass = 3
# The GPU partial segment is NOT in the final result (segments list reset),
# but on_segment was called live as segments streamed.
# This is acceptable — on_segment is a live progress callback.
# The important thing is that result.segments contains only CPU segments.
assert all(s.text.startswith(" Segment") for s in result.segments)
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_reports_missing_socksio_for_proxy(mock_model_cls):
mock_model_cls.side_effect = ImportError(
"Using SOCKS proxy, but the 'socksio' package is not installed."
) )
with pytest.raises(RuntimeError, match="socksio"): def backend_for_device(device, **kwargs):
return cuda_backend if device == "cuda" else cpu_backend
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_reports_missing_socksio_for_proxy(mock_get_backend):
backend = _make_backend(
create_model_error=ImportError(
"Using SOCKS proxy, but the 'socksio' package is not installed."
),
)
mock_get_backend.return_value = backend
# ImportError is not caught as backend error → propagates
with pytest.raises(ImportError, match="socksio"):
transcribe( transcribe(
file_path=Path("test.mp3"), file_path=Path("test.mp3"),
model_name="tiny", model_name="tiny",
@@ -235,14 +202,10 @@ def test_transcribe_reports_missing_socksio_for_proxy(mock_model_cls):
) )
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_reports_status_transitions(mock_model_cls): def test_transcribe_reports_status_transitions(mock_get_backend):
raw_segments = _make_raw_segments(1) backend = _make_backend(transcribe_result=_make_result(count=1))
info = _make_info() mock_get_backend.return_value = backend
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
mock_model_cls.return_value = instance
statuses: list[str] = [] statuses: list[str] = []
@@ -253,14 +216,138 @@ def test_transcribe_reports_status_transitions(mock_model_cls):
on_status=statuses.append, on_status=statuses.append,
) )
assert statuses == [ # load_model reports init status, _transcribe_file reports transcribe status
"Инициализирую модель на cpu...", assert any("Инициализирую модель" in s for s in statuses)
"Транскрибирую...", assert any("Транскрибирую" in s for s in statuses)
"Транскрибирую... 00:04 / 01:00 [1 сегм.]",
]
@patch("local_transcriber.transcriber.snapshot_download") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_strict_cuda_error(mock_get_backend):
"""strict_device=True + CUDA error -> raise, без fallback."""
backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
mock_get_backend.return_value = backend
with pytest.raises(RuntimeError, match="CUDA out of memory"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=True,
)
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_non_strict_cuda_fallback(mock_get_backend):
"""strict_device=False + CUDA error -> fallback на CPU."""
cuda_backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
cpu_backend = _make_backend(
transcribe_result=_make_result(count=2, device_used="cpu"),
model_path="/mock/cpu/model",
)
def backend_for_device(device, **kwargs):
return cuda_backend if device == "cuda" else cpu_backend
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=False,
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_strict_cuda_error_during_transcription(mock_get_backend):
"""strict_device=True + CUDA error during transcription -> raise."""
backend = _make_backend(
transcribe_error=RuntimeError("CUDA error during transcription"),
)
mock_get_backend.return_value = backend
with pytest.raises(RuntimeError, match="CUDA error during transcription"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=True,
)
# === load_model() tests ===
@patch("local_transcriber.transcriber.get_backend")
def test_load_model_cuda_fallback(mock_get_backend):
cuda_backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
cpu_model = MagicMock()
cpu_backend = _make_backend(model=cpu_model, model_path="/mock/cpu/model")
def backend_for_device(device, **kwargs):
return cuda_backend if device == "cuda" else cpu_backend
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
model, actual_device, backend, model_path = load_model("tiny", "cuda", "int8")
assert actual_device == "cpu"
assert model is cpu_model
@patch("local_transcriber.transcriber.get_backend")
def test_load_model_strict_raises(mock_get_backend):
backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
mock_get_backend.return_value = backend
with pytest.raises(RuntimeError, match="CUDA out of memory"):
load_model("tiny", "cuda", "int8", strict_device=True)
@patch("local_transcriber.transcriber.get_backend")
def test_load_model_returns_backend_and_path(mock_get_backend):
backend = _make_backend(model_path="/mock/model/path")
mock_get_backend.return_value = backend
model, actual_device, returned_backend, model_path = load_model("tiny", "cpu", "int8")
assert returned_backend is backend
assert model_path == "/mock/model/path"
assert actual_device == "cpu"
# === _transcribe_file() tests ===
def test__transcribe_file_basic():
result_data = _make_result(count=2)
backend = _make_backend(transcribe_result=result_data)
tfr = _transcribe_file(
model=MagicMock(),
actual_device="cpu",
backend=backend,
model_path="/mock/model",
file_path=Path("test.mp3"),
model_name="tiny",
compute_type="int8",
)
assert len(tfr.result.segments) == 2
assert tfr.actual_device == "cpu"
assert tfr.backend is backend
assert tfr.model_path == "/mock/model"
# === ensure_model_available() tests (через FasterWhisperBackend) ===
@patch("local_transcriber.backends.faster_whisper.snapshot_download")
def test_ensure_model_available_uses_cache_first(mock_snapshot_download, tmp_path): def test_ensure_model_available_uses_cache_first(mock_snapshot_download, tmp_path):
model_dir = _create_model_dir(tmp_path / "cache-model") model_dir = _create_model_dir(tmp_path / "cache-model")
mock_snapshot_download.return_value = str(model_dir) mock_snapshot_download.return_value = str(model_dir)
@@ -268,22 +355,15 @@ def test_ensure_model_available_uses_cache_first(mock_snapshot_download, tmp_pat
result = ensure_model_available("large-v3") result = ensure_model_available("large-v3")
assert result == str(model_dir) assert result == str(model_dir)
mock_snapshot_download.assert_called_once_with( mock_snapshot_download.assert_called_once()
"Systran/faster-whisper-large-v3", assert mock_snapshot_download.call_args.kwargs["local_files_only"] is True
local_files_only=True,
allow_patterns=[
"config.json",
"preprocessor_config.json",
"model.bin",
"tokenizer.json",
"vocabulary.*",
],
)
@patch("local_transcriber.transcriber._validate_model_dir") @patch("local_transcriber.backends.faster_whisper._validate_model_dir")
@patch("local_transcriber.transcriber.snapshot_download") @patch("local_transcriber.backends.faster_whisper.snapshot_download")
def test_ensure_model_available_downloads_on_cache_miss(mock_snapshot_download, mock_validate_model_dir): def test_ensure_model_available_downloads_on_cache_miss(mock_snapshot_download, mock_validate_model_dir):
from huggingface_hub.errors import LocalEntryNotFoundError
mock_snapshot_download.side_effect = [ mock_snapshot_download.side_effect = [
LocalEntryNotFoundError("not cached"), LocalEntryNotFoundError("not cached"),
"/downloaded/model", "/downloaded/model",
@@ -295,10 +375,8 @@ def test_ensure_model_available_downloads_on_cache_miss(mock_snapshot_download,
assert result == "/downloaded/model" assert result == "/downloaded/model"
assert mock_snapshot_download.call_args_list[0].kwargs["local_files_only"] is True assert mock_snapshot_download.call_args_list[0].kwargs["local_files_only"] is True
assert mock_snapshot_download.call_args_list[1].kwargs["local_files_only"] is False assert mock_snapshot_download.call_args_list[1].kwargs["local_files_only"] is False
assert statuses == [ assert "Проверяю кэш модели large-v3..." in statuses
"Проверяю кэш модели large-v3...", assert "Скачиваю модель large-v3 из Hugging Face..." in statuses
"Скачиваю модель large-v3 из Hugging Face...",
]
def test_ensure_model_available_accepts_local_directory(tmp_path): def test_ensure_model_available_accepts_local_directory(tmp_path):
@@ -311,7 +389,10 @@ def test_ensure_model_available_accepts_local_directory(tmp_path):
def test_ensure_model_available_accepts_repo_id(tmp_path): def test_ensure_model_available_accepts_repo_id(tmp_path):
model_dir = _create_model_dir(tmp_path / "repo-model") model_dir = _create_model_dir(tmp_path / "repo-model")
with patch("local_transcriber.transcriber.snapshot_download", return_value=str(model_dir)) as mock_snapshot_download: with patch(
"local_transcriber.backends.faster_whisper.snapshot_download",
return_value=str(model_dir),
) as mock_snapshot_download:
result = ensure_model_available("org/model") result = ensure_model_available("org/model")
assert result == str(model_dir) assert result == str(model_dir)
@@ -323,7 +404,7 @@ def test_ensure_model_available_rejects_unsupported_alias():
ensure_model_available("distil-large-v3") ensure_model_available("distil-large-v3")
@patch("local_transcriber.transcriber.snapshot_download") @patch("local_transcriber.backends.faster_whisper.snapshot_download")
def test_ensure_model_available_redownloads_incomplete_cache(mock_snapshot_download, tmp_path): def test_ensure_model_available_redownloads_incomplete_cache(mock_snapshot_download, tmp_path):
incomplete = tmp_path / "incomplete" incomplete = tmp_path / "incomplete"
incomplete.mkdir() incomplete.mkdir()
@@ -349,11 +430,7 @@ def test_ensure_model_available_redownloads_incomplete_cache(mock_snapshot_downl
result = ensure_model_available("large-v3", on_status=statuses.append) result = ensure_model_available("large-v3", on_status=statuses.append)
assert result == str(complete) assert result == str(complete)
assert statuses == [ assert "Кэш модели large-v3 неполный, докачиваю..." in statuses
"Проверяю кэш модели large-v3...",
"Кэш модели large-v3 неполный, докачиваю...",
"Скачиваю модель large-v3 из Hugging Face...",
]
def test_ensure_model_available_rejects_incomplete_local_directory(tmp_path): def test_ensure_model_available_rejects_incomplete_local_directory(tmp_path):
@@ -365,109 +442,103 @@ def test_ensure_model_available_rejects_incomplete_local_directory(tmp_path):
ensure_model_available(str(model_dir)) ensure_model_available(str(model_dir))
@patch("local_transcriber.transcriber.WhisperModel") # === Cross-backend fallback (openvino → cpu) ===
def test_transcribe_strict_cuda_error(mock_model_cls):
"""strict_device=True + CUDA error -> raise, без fallback."""
mock_model_cls.side_effect = RuntimeError("CUDA out of memory")
with pytest.raises(RuntimeError, match="CUDA out of memory"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=True,
)
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_non_strict_cuda_fallback(mock_model_cls): def test_load_model_openvino_fallback_to_cpu(mock_get_backend):
"""strict_device=False + CUDA error -> fallback на CPU.""" """OpenVINO ошибка при init → fallback на CPU (FasterWhisper)."""
raw_segments = _make_raw_segments(2) ov_backend = _make_backend(
info = _make_info() create_model_error=RuntimeError("OpenVINO model load failed"),
)
cpu_model = MagicMock()
cpu_backend = _make_backend(model=cpu_model, model_path="/mock/cpu/model")
cpu_instance = MagicMock() def backend_for_device(device, **kwargs):
cpu_instance.transcribe.return_value = (iter(raw_segments), info) return ov_backend if device == "openvino" else cpu_backend
def model_side_effect(model_name, device, compute_type): mock_get_backend.side_effect = backend_for_device
if device == "cuda":
raise RuntimeError("CUDA out of memory")
return cpu_instance
mock_model_cls.side_effect = model_side_effect
with pytest.warns(UserWarning, match="Переключение на CPU"): with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe( model, actual_device, backend, model_path = load_model(
file_path=Path("test.mp3"), "medium", "openvino", "int8",
model_name="tiny",
device="cuda",
strict_device=False,
) )
assert result.device_used == "cpu"
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_strict_cuda_error_during_transcription(mock_model_cls):
"""strict_device=True + CUDA error during transcription -> raise."""
cuda_instance = MagicMock()
cuda_instance.transcribe.side_effect = RuntimeError("CUDA error during transcription")
mock_model_cls.return_value = cuda_instance
with pytest.raises(RuntimeError, match="CUDA error during transcription"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=True,
)
# === load_model tests ===
@patch("local_transcriber.transcriber.WhisperModel")
def test_load_model_cuda_fallback(mock_model_cls):
cpu_instance = MagicMock()
def model_side_effect(model_name, device, compute_type):
if device == "cuda":
raise RuntimeError("CUDA out of memory")
return cpu_instance
mock_model_cls.side_effect = model_side_effect
with pytest.warns(UserWarning, match="Переключение на CPU"):
model, actual_device = load_model("tiny", "cuda", "int8")
assert actual_device == "cpu" assert actual_device == "cpu"
assert model is cpu_instance assert model is cpu_model
assert backend is cpu_backend
assert model_path == "/mock/cpu/model"
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_load_model_strict_raises(mock_model_cls): def test_transcribe_file_openvino_midstream_fallback(mock_get_backend):
mock_model_cls.side_effect = RuntimeError("CUDA out of memory") """OpenVINO ошибка при транскрипции → fallback на CPU."""
ov_backend = _make_backend(
with pytest.raises(RuntimeError, match="CUDA out of memory"): transcribe_error=RuntimeError("OpenVINO inference error"),
load_model("tiny", "cuda", "int8", strict_device=True) )
cpu_backend = _make_backend(
transcribe_result=_make_result(count=2, device_used="cpu"),
@patch("local_transcriber.transcriber.WhisperModel") model_path="/mock/cpu/model",
def test__transcribe_file_basic(mock_model_cls):
raw_segments = _make_raw_segments(2)
info = _make_info()
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
tfr = _transcribe_file(
model=instance,
actual_device="cpu",
file_path=Path("test.mp3"),
model_name="tiny",
compute_type="int8",
) )
assert len(tfr.result.segments) == 2 def backend_for_device(device, **kwargs):
return ov_backend if device == "openvino" else cpu_backend
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
tfr = _transcribe_file(
model=MagicMock(),
actual_device="openvino",
backend=ov_backend,
model_path="/mock/ov/model",
file_path=Path("test.mp3"),
model_name="medium",
compute_type="int8",
)
assert tfr.actual_device == "cpu" assert tfr.actual_device == "cpu"
assert tfr.model is instance assert tfr.backend is cpu_backend
assert tfr.model_path == "/mock/cpu/model"
@patch("local_transcriber.transcriber.get_backend")
def test_openvino_strict_device_no_fallback(mock_get_backend):
"""strict_device=True + OpenVINO ошибка → raise."""
backend = _make_backend(
create_model_error=RuntimeError("OpenVINO model load failed"),
)
mock_get_backend.return_value = backend
with pytest.raises(RuntimeError, match="OpenVINO"):
load_model("medium", "openvino", "int8", strict_device=True)
@patch("local_transcriber.transcriber.get_backend")
def test_openvino_runtime_error_triggers_fallback(mock_get_backend):
"""Любой RuntimeError от OpenVINO бэкенда → fallback."""
ov_backend = _make_backend(
create_model_error=RuntimeError("Exception from src/inference/..."),
)
cpu_backend = _make_backend(model_path="/mock/cpu/model")
def backend_for_device(device, **kwargs):
return ov_backend if device == "openvino" else cpu_backend
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
_, actual_device, _, _ = load_model("medium", "openvino", "int8")
assert actual_device == "cpu"
def test_ensure_model_available_openvino_default_compute_type():
"""ensure_model_available(device='openvino') без compute_type не падает."""
from local_transcriber.backends.openvino import OpenVINOBackend
backend = OpenVINOBackend(compute_type_explicit=True)
# Проверяем что _resolve_repo работает с дефолтным compute_type для openvino (int8)
repo, ct = backend._resolve_repo("medium", "int8")
assert repo == "OpenVINO/whisper-medium-int8-ov"
assert ct == "int8"
+28
View File
@@ -58,6 +58,34 @@ def test_build_output_path_custom():
def test_detect_device_explicit(): def test_detect_device_explicit():
assert detect_device("cpu") == "cpu" assert detect_device("cpu") == "cpu"
assert detect_device("cuda") == "cuda" assert detect_device("cuda") == "cuda"
assert detect_device("openvino") == "openvino"
def test_detect_device_auto_openvino():
"""Нет nvidia-smi, есть openvino_genai, x86_64 → openvino."""
with (
patch("local_transcriber.utils.shutil.which", return_value=None),
patch("local_transcriber.utils._is_openvino_available", return_value=True),
):
assert detect_device("auto") == "openvino"
def test_detect_device_cuda_over_openvino():
"""nvidia-smi доступен и openvino тоже → cuda побеждает."""
with (
patch("local_transcriber.utils.shutil.which", return_value="/usr/bin/nvidia-smi"),
patch("local_transcriber.utils._is_openvino_available", return_value=True),
):
assert detect_device("auto") == "cuda"
def test_detect_device_auto_cpu_fallback():
"""Ни nvidia-smi, ни openvino → cpu."""
with (
patch("local_transcriber.utils.shutil.which", return_value=None),
patch("local_transcriber.utils._is_openvino_available", return_value=False),
):
assert detect_device("auto") == "cpu"
def test_get_gpu_name_no_nvidia_smi(): def test_get_gpu_name_no_nvidia_smi():