refactor(transcriber): введена pluggable-архитектура бэкендов транскрипции

- Зачем:
  - подготовка к добавлению OpenVINO бэкенда для ускорения на x86 CPU без CUDA.
  - архитектура должна позволять добавлять новые бэкенды (CoreML, AMD XDNA) без переписывания кода.
- Что:
  - создан types.py с общими типами (Segment, TranscribeResult, TranscribeFileResult).
  - создан backends/base.py с Backend Protocol (3 метода: ensure_model_available, create_model, transcribe).
  - создан backends/faster_whisper.py — текущий код вынесен из transcriber.py в FasterWhisperBackend.
  - transcriber.py переделан в оркестратор: load_model() владеет полным пайплайном (ensure + create), CLI больше не вызывает ensure_model_available() отдельно.
  - TranscribeFileResult расширен полями backend и model_path для корректного cross-backend fallback в батч-режиме.
  - device_used проставляется оркестратором, а не бэкендом.
  - cli.py: вынесен _format_device_info(), подготовлен к openvino.
  - тесты обновлены: mock-точки перенесены с WhisperModel на get_backend/бэкенд-объекты.
- Проверка:
  - uv run pytest -v — 98 passed.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-03-21 23:18:39 +03:00
co-authored by Claude Opus 4.6
parent 8632960354
commit 4c5399c2fc
9 changed files with 747 additions and 610 deletions
@@ -0,0 +1,28 @@
"""Реестр бэкендов транскрипции и выбор бэкенда по устройству."""
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from .base import Backend
def get_backend(device: str) -> Backend:
"""Возвращает экземпляр бэкенда для указанного устройства.
Импорты ленивые — бэкенд загружается только при запросе.
"""
if device == "openvino":
try:
from .openvino import OpenVINOBackend
except ImportError:
raise ValueError(
"OpenVINO бэкенд недоступен. Установите: pip install openvino-genai"
) from None
return OpenVINOBackend()
# cuda, cpu и всё остальное → faster-whisper
from .faster_whisper import FasterWhisperBackend
return FasterWhisperBackend()
+46
View File
@@ -0,0 +1,46 @@
"""Протокол бэкенда транскрипции."""
from __future__ import annotations
from collections.abc import Callable
from pathlib import Path
from typing import Any, Protocol
from local_transcriber.types import Segment, TranscribeResult
class Backend(Protocol):
"""Минимальный интерфейс бэкенда транскрипции.
Бэкенды реализуют этот протокол (structural typing) —
наследование не требуется.
"""
def ensure_model_available(
self,
model_name: str,
compute_type: str,
on_status: Callable[[str], None] | None = None,
) -> str:
"""Гарантирует наличие модели, возвращает путь к файлам."""
...
def create_model(
self,
model_path: str,
device: str,
compute_type: str,
) -> Any:
"""Создаёт модель. Возвращает backend-специфичный объект."""
...
def transcribe(
self,
model: Any,
file_path: Path,
language: str | None,
on_segment: Callable[[Segment], None] | None = None,
on_status: Callable[[str], None] | None = None,
) -> TranscribeResult:
"""Транскрибирует файл, возвращает результат."""
...
@@ -0,0 +1,181 @@
"""Бэкенд транскрипции на основе 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 ensure_model_available(
self,
model_name: str,
compute_type: str,
on_status: Callable[[str], None] | None = None,
) -> str:
"""Резолвит alias модели в repo_id и гарантирует наличие файлов."""
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
+30 -31
View File
@@ -14,9 +14,7 @@ from .transcriber import (
Segment, Segment,
_is_cuda_error, _is_cuda_error,
_transcribe_file, _transcribe_file,
ensure_model_available,
load_model, load_model,
transcribe,
) )
from .utils import ( from .utils import (
build_output_path, build_output_path,
@@ -31,6 +29,16 @@ app = typer.Typer()
console = Console(stderr=True) console = Console(stderr=True)
def _format_device_info(device_used: str) -> str:
"""Формирует строку устройства для шапки транскрипта."""
if device_used == "cuda":
gpu_name = get_gpu_name()
return f"CUDA ({gpu_name or 'Unknown GPU'})"
if device_used == "openvino":
return "OpenVINO (CPU)"
return "CPU"
@app.command() @app.command()
def main( def main(
files: list[Path] = typer.Argument(..., help="Пути к аудио/видеофайлам"), files: list[Path] = typer.Argument(..., help="Пути к аудио/видеофайлам"),
@@ -42,7 +50,8 @@ 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) [по умолч.: 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,
@@ -120,7 +129,6 @@ 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)
@@ -131,15 +139,11 @@ def _run_single(
f"Compute: [bold]{defaults['compute_type']}[/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,
) )
@@ -147,8 +151,10 @@ def _run_single(
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 +182,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(
@@ -232,15 +233,12 @@ 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,
) )
@@ -277,8 +275,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 +291,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 +304,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,
+1 -1
View File
@@ -4,7 +4,7 @@ from dataclasses import dataclass
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
from .transcriber import Segment, TranscribeResult from .types import Segment, TranscribeResult
_PAUSE_THRESHOLD_S = 2.0 # пауза между сегментами для разбиения на абзацы _PAUSE_THRESHOLD_S = 2.0 # пауза между сегментами для разбиения на абзацы
_MAX_PARAGRAPH_S = 60.0 # максимальная длительность абзаца _MAX_PARAGRAPH_S = 60.0 # максимальная длительность абзаца
+62 -191
View File
@@ -1,65 +1,18 @@
"""Обёртка над faster-whisper: загрузка моделей, транскрипция, CUDA fallback.""" """Оркестрация транскрипции: выбор бэкенда, загрузка модели, fallback."""
import warnings import warnings
from collections.abc import Callable from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import Any
# Должен быть ДО импорта faster_whisper / ctranslate2 from local_transcriber.backends import get_backend
from local_transcriber._cuda_bootstrap import ensure_cublas_loadable
ensure_cublas_loadable() # Re-export из types.py для обратной совместимости
from local_transcriber.types import ( # noqa: F401
from faster_whisper import WhisperModel # noqa: E402 Segment,
from huggingface_hub import snapshot_download TranscribeFileResult,
from huggingface_hub.errors import LocalEntryNotFoundError TranscribeResult,
)
MODEL_REPOS = {
"tiny": "Systran/faster-whisper-tiny",
"base": "Systran/faster-whisper-base",
"small": "Systran/faster-whisper-small",
"medium": "Systran/faster-whisper-medium",
"large-v3": "Systran/faster-whisper-large-v3",
}
# allow — фильтр для snapshot_download (какие файлы скачивать из репозитория);
# required — для валидации (что обязано быть после скачивания/в локальной модели)
MODEL_ALLOW_PATTERNS = [
"config.json",
"preprocessor_config.json",
"model.bin",
"tokenizer.json",
"vocabulary.*",
]
MODEL_REQUIRED_FILES = [
"config.json",
"model.bin",
"tokenizer.json",
]
@dataclass
class Segment:
start: float # seconds
end: float # seconds
text: str
@dataclass
class TranscribeResult:
segments: list[Segment]
language: str
language_probability: float
duration: float # seconds
device_used: str # "cpu" / "cuda"
@dataclass
class TranscribeFileResult:
result: TranscribeResult
model: WhisperModel
actual_device: str
def load_model( def load_model(
@@ -68,15 +21,21 @@ 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]: ) -> tuple[Any, str, Any, str]:
"""Загружает модель с CUDA-фолбеком. Возвращает (model, actual_device).""" """Загружает модель: ensure + create с fallback.
Возвращает (model, actual_device, backend, model_path).
"""
backend = get_backend(device)
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 +44,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 +67,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 +114,12 @@ 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,
)
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 +127,30 @@ def transcribe(
def ensure_model_available( def ensure_model_available(
model_name: str, model_name: str,
device: str = "cpu",
compute_type: str = "float32",
on_status: Callable[[str], None] | None = None, on_status: Callable[[str], None] | None = None,
) -> str: ) -> str:
"""Резолвит alias модели в repo_id и гарантирует наличие файлов. """Публичный helper: гарантирует наличие модели для указанного бэкенда."""
backend = get_backend(device)
Стратегия: cache-first (``local_files_only=True``), затем download. return backend.ensure_model_available(model_name, compute_type, on_status)
Два вызова ``snapshot_download`` — чтобы не лезть в сеть, если модель уже в кэше.
"""
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_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)
# openvino и другие бэкенды: конкретные паттерны ошибок добавим
def _fmt_time(seconds: float) -> str: # при реализации бэкенда; пока — не маскируем ошибки
m, s = divmod(int(seconds), 60) return False
h, m = divmod(m, 60)
return f"{h}:{m:02d}:{s:02d}" if h else f"{m:02d}:{s:02d}"
def _notify_status(on_status: Callable[[str], None] | None, message: str) -> None: def _notify_status(on_status: Callable[[str], None] | None, message: str) -> None:
if on_status is not None: if on_status is not None:
on_status(message) on_status(message)
def _resolve_model_repo(model_name: str) -> str:
if "/" in model_name:
return model_name
repo_id = MODEL_REPOS.get(model_name)
if repo_id is None:
expected = ", ".join(MODEL_REPOS)
raise ValueError(f"Неподдерживаемая модель '{model_name}'. Ожидалось одно из: {expected}")
return repo_id
def _snapshot_download(repo_id: str, local_files_only: bool) -> str:
try:
return snapshot_download(
repo_id,
local_files_only=local_files_only,
allow_patterns=MODEL_ALLOW_PATTERNS,
)
except ImportError as exc:
if _is_missing_socksio_error(exc):
raise RuntimeError(
"Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, "
"нужная для загрузки модели из Hugging Face через proxy. "
"Обновите окружение: `uv sync`."
) from exc
raise
def _validate_model_dir(model_dir: Path) -> None:
missing = [
filename for filename in MODEL_REQUIRED_FILES if not (model_dir / filename).exists()
]
if not any(model_dir.glob("vocabulary.*")):
missing.append("vocabulary.*")
if missing:
missing_str = ", ".join(missing)
raise ValueError(f"Неполная локальная модель в '{model_dir}': отсутствуют {missing_str}")
+34
View File
@@ -0,0 +1,34 @@
"""Общие типы данных для всех бэкендов транскрипции."""
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
@dataclass
class Segment:
start: float # seconds
end: float # seconds
text: str
@dataclass
class TranscribeResult:
segments: list[Segment]
language: str
language_probability: float
duration: float # seconds
device_used: str # "cpu" / "cuda" / "openvino"
@dataclass
class TranscribeFileResult:
result: TranscribeResult
model: Any # backend-specific model handle
actual_device: str
backend: Any = None # backend instance (для переиспользования в батче)
model_path: str = "" # путь к модели (меняется при cross-backend fallback)
StatusCallback = Callable[[str], None] | None
SegmentCallback = Callable[[Segment], None] | None
+103 -94
View File
@@ -24,12 +24,21 @@ def _make_model():
return MagicMock(name="WhisperModel") return MagicMock(name="WhisperModel")
def _make_tfr(result=None, model=None, actual_device="cpu"): def _make_backend():
return MagicMock(name="Backend")
def _make_tfr(result=None, model=None, actual_device="cpu", backend=None, model_path="/models/medium"):
if result is None: if result is None:
result = _make_result() result = _make_result()
if model is None: if model is None:
model = _make_model() model = _make_model()
return TranscribeFileResult(result=result, model=model, actual_device=actual_device) if backend is None:
backend = _make_backend()
return TranscribeFileResult(
result=result, model=model, actual_device=actual_device,
backend=backend, model_path=model_path,
)
def _single_patches(result=None, tmp_file=None, actual_device="cpu"): def _single_patches(result=None, tmp_file=None, actual_device="cpu"):
@@ -37,13 +46,13 @@ def _single_patches(result=None, tmp_file=None, actual_device="cpu"):
if result is None: if result is None:
result = _make_result(device_used=actual_device) result = _make_result(device_used=actual_device)
model = _make_model() model = _make_model()
tfr = TranscribeFileResult(result=result, model=model, actual_device=actual_device) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device=actual_device, backend=backend)
return [ return [
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=tmp_file), patch("local_transcriber.cli.validate_input_file", return_value=tmp_file),
patch("local_transcriber.cli.detect_device", return_value=actual_device), patch("local_transcriber.cli.detect_device", return_value=actual_device),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, actual_device, backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, actual_device)),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
] ]
@@ -54,7 +63,7 @@ def test_cli_happy_path_exit_code_zero(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
patches = _single_patches(tmp_file=audio) patches = _single_patches(tmp_file=audio)
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6]: with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5]:
out = runner.invoke(app, [str(audio)]) out = runner.invoke(app, [str(audio)])
assert out.exit_code == 0 assert out.exit_code == 0
@@ -65,22 +74,22 @@ def test_cli_default_options_passed_to_transcribe(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
runner.invoke(app, [str(audio)]) runner.invoke(app, [str(audio)])
call_kwargs = mock_transcribe_file.call_args[1] call_kwargs = mock_transcribe_file.call_args[1]
assert call_kwargs["model_name"] == "/models/medium" assert call_kwargs["model_name"] == "medium"
assert call_kwargs["compute_type"] == "float32" assert call_kwargs["compute_type"] == "float32"
assert call_kwargs["language"] == "ru" assert call_kwargs["language"] == "ru"
assert call_kwargs["on_segment"] is None # verbose=False assert call_kwargs["on_segment"] is None # verbose=False
@@ -91,15 +100,15 @@ def test_cli_custom_options(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result(device_used="cuda") result = _make_result(device_used="cuda")
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model, actual_device="cuda") backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device="cuda", backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/small"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/small")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"), patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"),
@@ -113,7 +122,7 @@ def test_cli_custom_options(tmp_path):
]) ])
call_kwargs = mock_transcribe_file.call_args[1] call_kwargs = mock_transcribe_file.call_args[1]
assert call_kwargs["model_name"] == "/models/small" assert call_kwargs["model_name"] == "small"
assert call_kwargs["language"] == "ru" assert call_kwargs["language"] == "ru"
assert call_kwargs["compute_type"] == "float16" assert call_kwargs["compute_type"] == "float16"
@@ -123,15 +132,15 @@ def test_cli_verbose_passes_on_segment_callback(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -148,7 +157,7 @@ def test_cli_empty_speech_warning(tmp_path):
result = _make_result(segments=[]) result = _make_result(segments=[])
patches = _single_patches(result=result, tmp_file=audio) patches = _single_patches(result=result, tmp_file=audio)
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6]: with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5]:
out = runner.invoke(app, [str(audio)]) out = runner.invoke(app, [str(audio)])
assert out.exit_code == 0 assert out.exit_code == 0
@@ -162,14 +171,14 @@ def test_cli_default_output_path(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript", mock_write), patch("local_transcriber.cli.write_transcript", mock_write),
): ):
@@ -187,14 +196,14 @@ def test_cli_custom_output_path(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript", mock_write), patch("local_transcriber.cli.write_transcript", mock_write),
): ):
@@ -209,15 +218,15 @@ def test_cli_passes_status_callback_to_transcribe(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -228,28 +237,27 @@ def test_cli_passes_status_callback_to_transcribe(tmp_path):
assert callable(call_kwargs["on_status"]) assert callable(call_kwargs["on_status"])
def test_cli_resolves_model_before_transcribe(tmp_path): def test_cli_load_model_called_with_model_name(tmp_path):
"""load_model receives model name from defaults, handles ensure internally."""
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
mock_transcribe_file = MagicMock(return_value=tfr) tfr = _make_tfr(result=result, model=model, backend=backend)
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/large-v3"))
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3") as mock_ensure, patch("local_transcriber.cli.load_model", mock_load_model),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
runner.invoke(app, [str(audio), "--model", "large-v3"]) runner.invoke(app, [str(audio), "--model", "large-v3"])
mock_ensure.assert_called_once() assert mock_load_model.call_args[0][0] == "large-v3"
call_kwargs = mock_transcribe_file.call_args[1]
assert call_kwargs["model_name"] == "/models/large-v3"
def test_cli_windows_cuda_diagnostic(tmp_path): def test_cli_windows_cuda_diagnostic(tmp_path):
@@ -257,13 +265,13 @@ def test_cli_windows_cuda_diagnostic(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")),
patch("local_transcriber.cli.sys") as mock_sys, patch("local_transcriber.cli.sys") as mock_sys,
): ):
@@ -280,13 +288,13 @@ def test_cli_linux_cuda_error_no_windows_hint(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")),
patch("local_transcriber.cli.sys") as mock_sys, patch("local_transcriber.cli.sys") as mock_sys,
): ):
@@ -303,14 +311,14 @@ def test_cli_device_fallback_warning(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result(device_used="cpu") result = _make_result(device_used="cpu")
model = _make_model() model = _make_model()
tfr = TranscribeFileResult(result=result, model=model, actual_device="cpu") backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device="cpu", backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -325,15 +333,15 @@ def test_cli_strict_device_passed_to_transcribe(tmp_path):
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
result = _make_result(device_used="cuda") result = _make_result(device_used="cuda")
model = _make_model() model = _make_model()
tfr = TranscribeFileResult(result=result, model=model, actual_device="cuda") backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device="cuda", backend=backend)
mock_transcribe_file = MagicMock(return_value=tfr) mock_transcribe_file = MagicMock(return_value=tfr)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"), patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"),
@@ -344,15 +352,14 @@ def test_cli_strict_device_passed_to_transcribe(tmp_path):
mock_transcribe_file.reset_mock() mock_transcribe_file.reset_mock()
result_cpu = _make_result(device_used="cpu") result_cpu = _make_result(device_used="cpu")
tfr_cpu = TranscribeFileResult(result=result_cpu, model=model, actual_device="cpu") tfr_cpu = _make_tfr(result=result_cpu, model=model, backend=backend)
mock_transcribe_file.return_value = tfr_cpu mock_transcribe_file.return_value = tfr_cpu
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -366,13 +373,13 @@ def test_cli_keyboard_interrupt(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=KeyboardInterrupt), patch("local_transcriber.cli._transcribe_file", side_effect=KeyboardInterrupt),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -399,13 +406,13 @@ def test_cli_unexpected_error_verbose_traceback(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -420,13 +427,13 @@ def test_cli_unexpected_error_no_verbose_hint(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -448,14 +455,14 @@ def test_cli_batch_two_files(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -470,19 +477,18 @@ def test_cli_batch_skips_existing(tmp_path):
b = tmp_path / "b.mp3" b = tmp_path / "b.mp3"
a.write_bytes(b"fake") a.write_bytes(b"fake")
b.write_bytes(b"fake") b.write_bytes(b"fake")
# Create transcript for a
(tmp_path / "a-transcript.md").write_text("existing") (tmp_path / "a-transcript.md").write_text("existing")
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -524,14 +530,14 @@ def test_cli_batch_force_overwrites(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -550,7 +556,8 @@ def test_cli_batch_per_file_error(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
call_count = 0 call_count = 0
def transcribe_side_effect(**kwargs): def transcribe_side_effect(**kwargs):
@@ -564,8 +571,7 @@ def test_cli_batch_per_file_error(tmp_path):
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=transcribe_side_effect), patch("local_transcriber.cli._transcribe_file", side_effect=transcribe_side_effect),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -584,7 +590,8 @@ def test_cli_batch_invalid_in_prescan(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
tfr = _make_tfr(result=result, model=model, backend=backend)
def validate_side_effect(p): def validate_side_effect(p):
if not p.exists(): if not p.exists():
@@ -595,8 +602,7 @@ def test_cli_batch_invalid_in_prescan(tmp_path):
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=validate_side_effect), patch("local_transcriber.cli.validate_input_file", side_effect=validate_side_effect),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -634,42 +640,46 @@ def test_cli_config_applied(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
result = _make_result() result = _make_result()
tfr = _make_tfr(result=result, model=model) tfr = _make_tfr(result=result, model=model, backend=backend)
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/tiny"))
with ( with (
patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}), patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/tiny") as mock_ensure, patch("local_transcriber.cli.load_model", mock_load_model),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
runner.invoke(app, [str(audio)]) runner.invoke(app, [str(audio)])
mock_ensure.assert_called_once_with("tiny", on_status=mock_ensure.call_args[1]["on_status"]) # load_model receives model name from config
assert mock_load_model.call_args[0][0] == "tiny"
def test_cli_cli_overrides_config(tmp_path): def test_cli_cli_overrides_config(tmp_path):
audio = tmp_path / "test.mp3" audio = tmp_path / "test.mp3"
audio.write_bytes(b"fake") audio.write_bytes(b"fake")
model = _make_model() model = _make_model()
backend = _make_backend()
result = _make_result() result = _make_result()
tfr = _make_tfr(result=result, model=model) tfr = _make_tfr(result=result, model=model, backend=backend)
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/small"))
with ( with (
patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}), patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}),
patch("local_transcriber.cli.validate_input_file", return_value=audio), patch("local_transcriber.cli.validate_input_file", return_value=audio),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/small") as mock_ensure, patch("local_transcriber.cli.load_model", mock_load_model),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
runner.invoke(app, [str(audio), "--model", "small"]) runner.invoke(app, [str(audio), "--model", "small"])
mock_ensure.assert_called_once_with("small", on_status=mock_ensure.call_args[1]["on_status"]) # CLI --model overrides config
assert mock_load_model.call_args[0][0] == "small"
def test_cli_batch_fallback_warning(tmp_path): def test_cli_batch_fallback_warning(tmp_path):
@@ -681,14 +691,14 @@ def test_cli_batch_fallback_warning(tmp_path):
result = _make_result(device_used="cpu") result = _make_result(device_used="cpu")
model = _make_model() model = _make_model()
tfr = TranscribeFileResult(result=result, model=model, actual_device="cpu") backend = _make_backend()
tfr = _make_tfr(result=result, model=model, actual_device="cpu", backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -707,15 +717,15 @@ def test_cli_batch_empty_speech_warning(tmp_path):
result_empty = _make_result(segments=[]) result_empty = _make_result(segments=[])
result_ok = _make_result() result_ok = _make_result()
model = _make_model() model = _make_model()
tfr_empty = _make_tfr(result=result_empty, model=model) backend = _make_backend()
tfr_ok = _make_tfr(result=result_ok, model=model) tfr_empty = _make_tfr(result=result_empty, model=model, backend=backend)
tfr_ok = _make_tfr(result=result_ok, model=model, backend=backend)
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_empty, tfr_ok]), patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_empty, tfr_ok]),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -735,17 +745,16 @@ def test_cli_batch_midstream_fallback_warning(tmp_path):
model_gpu = _make_model() model_gpu = _make_model()
model_cpu = _make_model() model_cpu = _make_model()
backend = _make_backend()
result = _make_result(device_used="cpu") result = _make_result(device_used="cpu")
# First file triggers mid-stream fallback tfr_fallback = _make_tfr(result=result, model=model_cpu, actual_device="cpu", backend=backend)
tfr_fallback = TranscribeFileResult(result=result, model=model_cpu, actual_device="cpu") tfr_ok = _make_tfr(result=result, model=model_cpu, actual_device="cpu", backend=backend)
tfr_ok = TranscribeFileResult(result=result, model=model_cpu, actual_device="cpu")
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cuda"), patch("local_transcriber.cli.detect_device", return_value="cuda"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda", backend, "/models/medium")),
patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda")),
patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_fallback, tfr_ok]), patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_fallback, tfr_ok]),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
): ):
@@ -763,14 +772,14 @@ def test_cli_batch_model_loaded_once(tmp_path):
result = _make_result() result = _make_result()
model = _make_model() model = _make_model()
tfr = _make_tfr(result=result, model=model) backend = _make_backend()
mock_load_model = MagicMock(return_value=(model, "cpu")) tfr = _make_tfr(result=result, model=model, backend=backend)
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/medium"))
with ( with (
patch("local_transcriber.cli.load_config", return_value={}), patch("local_transcriber.cli.load_config", return_value={}),
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
patch("local_transcriber.cli.detect_device", return_value="cpu"), patch("local_transcriber.cli.detect_device", return_value="cpu"),
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
patch("local_transcriber.cli.load_model", mock_load_model), patch("local_transcriber.cli.load_model", mock_load_model),
patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli._transcribe_file", return_value=tfr),
patch("local_transcriber.cli.write_transcript"), patch("local_transcriber.cli.write_transcript"),
+262 -293
View File
@@ -1,9 +1,7 @@
from collections.abc import Generator
from pathlib import Path from pathlib import Path
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import pytest import pytest
from huggingface_hub.errors import LocalEntryNotFoundError
from local_transcriber.transcriber import ( from local_transcriber.transcriber import (
Segment, Segment,
@@ -15,24 +13,53 @@ from local_transcriber.transcriber import (
) )
def _make_raw_segments(count: int) -> list: # === Helpers ===
"""Create mock raw segments as returned by faster-whisper."""
segments = []
for i in range(count):
seg = MagicMock()
seg.start = float(i * 5)
seg.end = float(i * 5 + 4)
seg.text = f" Segment {i}"
segments.append(seg)
return segments
def _make_info(language: str = "ru", probability: float = 0.95, duration: float = 60.0): def _make_result(
info = MagicMock() count: int = 2,
info.language = language language: str = "ru",
info.language_probability = probability probability: float = 0.95,
info.duration = duration duration: float = 60.0,
return info device_used: str = "cpu",
) -> TranscribeResult:
segments = [
Segment(start=float(i * 5), end=float(i * 5 + 4), text=f" Segment {i}")
for i in range(count)
]
return TranscribeResult(
segments=segments,
language=language,
language_probability=probability,
duration=duration,
device_used=device_used,
)
def _make_backend(
model=None,
transcribe_result=None,
create_model_error=None,
transcribe_error=None,
model_path="/mock/model",
):
"""Создаёт mock-бэкенд с настраиваемым поведением."""
backend = MagicMock()
backend.ensure_model_available.return_value = model_path
if create_model_error:
backend.create_model.side_effect = create_model_error
else:
backend.create_model.return_value = model or MagicMock()
if transcribe_error:
backend.transcribe.side_effect = transcribe_error
elif transcribe_result:
backend.transcribe.return_value = transcribe_result
else:
backend.transcribe.return_value = _make_result()
return backend
def _create_model_dir(path: Path) -> Path: def _create_model_dir(path: Path) -> Path:
@@ -45,14 +72,14 @@ def _create_model_dir(path: Path) -> Path:
return path return path
@patch("local_transcriber.transcriber.WhisperModel") # === transcribe() tests ===
def test_transcribe_collects_segments(mock_model_cls):
raw_segments = _make_raw_segments(3)
info = _make_info()
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info) @patch("local_transcriber.transcriber.get_backend")
mock_model_cls.return_value = instance def test_transcribe_collects_segments(mock_get_backend):
result_data = _make_result(count=3)
backend = _make_backend(transcribe_result=result_data)
mock_get_backend.return_value = backend
result = transcribe( result = transcribe(
file_path=Path("test.mp3"), file_path=Path("test.mp3"),
@@ -68,14 +95,11 @@ def test_transcribe_collects_segments(mock_model_cls):
assert result.duration == 60.0 assert result.duration == 60.0
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_calls_on_segment(mock_model_cls): def test_transcribe_calls_on_segment(mock_get_backend):
raw_segments = _make_raw_segments(3) result_data = _make_result(count=3)
info = _make_info() backend = _make_backend(transcribe_result=result_data)
mock_get_backend.return_value = backend
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
mock_model_cls.return_value = instance
callback = MagicMock() callback = MagicMock()
@@ -86,28 +110,24 @@ def test_transcribe_calls_on_segment(mock_model_cls):
on_segment=callback, on_segment=callback,
) )
assert callback.call_count == 3 # on_segment is passed through to backend.transcribe
# Each call should receive a Segment instance call_args = backend.transcribe.call_args
for call_args in callback.call_args_list: assert call_args.kwargs.get("on_segment") is callback or call_args[0][3] is callback
seg = call_args[0][0]
assert isinstance(seg, Segment)
@patch("local_transcriber.transcriber.WhisperModel") @patch("local_transcriber.transcriber.get_backend")
def test_transcribe_cuda_fallback(mock_model_cls): def test_transcribe_cuda_fallback(mock_get_backend):
raw_segments = _make_raw_segments(2) """CUDA error at init -> fallback на CPU."""
info = _make_info() cuda_backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
cpu_backend = _make_backend(
transcribe_result=_make_result(count=2, device_used="cpu"),
model_path="/mock/cpu/model",
)
# First call (cuda) raises, second call (cpu) succeeds def backend_for_device(device):
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):
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):
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):
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):
@@ -363,111 +440,3 @@ def test_ensure_model_available_rejects_incomplete_local_directory(tmp_path):
with pytest.raises(ValueError, match="Неполная локальная модель"): with pytest.raises(ValueError, match="Неполная локальная модель"):
ensure_model_available(str(model_dir)) ensure_model_available(str(model_dir))
@patch("local_transcriber.transcriber.WhisperModel")
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")
def test_transcribe_non_strict_cuda_fallback(mock_model_cls):
"""strict_device=False + CUDA error -> fallback на CPU."""
raw_segments = _make_raw_segments(2)
info = _make_info()
cpu_instance = MagicMock()
cpu_instance.transcribe.return_value = (iter(raw_segments), info)
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"):
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.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 model is cpu_instance
@patch("local_transcriber.transcriber.WhisperModel")
def test_load_model_strict_raises(mock_model_cls):
mock_model_cls.side_effect = RuntimeError("CUDA out of memory")
with pytest.raises(RuntimeError, match="CUDA out of memory"):
load_model("tiny", "cuda", "int8", strict_device=True)
@patch("local_transcriber.transcriber.WhisperModel")
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
assert tfr.actual_device == "cpu"
assert tfr.model is instance