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:
@@ -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>
|
||||||
|
|||||||
@@ -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
@@ -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
|
|
||||||
```
|
```
|
||||||
После установки перезапустите терминал.
|
После установки перезапустите терминал.
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -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)}"
|
||||||
|
)
|
||||||
@@ -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,
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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 # максимальная длительность абзаца
|
||||||
|
|||||||
@@ -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}")
|
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
@@ -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"),
|
||||||
|
|||||||
@@ -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
@@ -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"
|
||||||
|
|||||||
@@ -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():
|
||||||
|
|||||||
Reference in New Issue
Block a user