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:
@@ -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()
|
||||
@@ -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
|
||||
@@ -14,9 +14,7 @@ from .transcriber import (
|
||||
Segment,
|
||||
_is_cuda_error,
|
||||
_transcribe_file,
|
||||
ensure_model_available,
|
||||
load_model,
|
||||
transcribe,
|
||||
)
|
||||
from .utils import (
|
||||
build_output_path,
|
||||
@@ -31,6 +29,16 @@ app = typer.Typer()
|
||||
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()
|
||||
def main(
|
||||
files: list[Path] = typer.Argument(..., help="Пути к аудио/видеофайлам"),
|
||||
@@ -42,7 +50,8 @@ def main(
|
||||
),
|
||||
output: Path | None = typer.Option(None, "--output", "-o", help="Путь к выходному файлу"),
|
||||
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(
|
||||
None, "--compute-type", show_default=False,
|
||||
@@ -120,7 +129,6 @@ def _run_single(
|
||||
validated_file = validate_input_file(file)
|
||||
requested_device = defaults["device"]
|
||||
resolved_device = detect_device(requested_device)
|
||||
# Если пользователь явно указал устройство — запрещаем fallback на CPU
|
||||
strict = requested_device != "auto"
|
||||
output_path = build_output_path(validated_file, output)
|
||||
|
||||
@@ -131,15 +139,11 @@ def _run_single(
|
||||
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:
|
||||
console.print(f" [{seg.start:.2f}s] {seg.text.strip()}")
|
||||
|
||||
model_obj, actual_device = load_model(
|
||||
model_path, resolved_device, defaults["compute_type"],
|
||||
model_obj, actual_device, backend, model_path = load_model(
|
||||
defaults["model"], resolved_device, defaults["compute_type"],
|
||||
on_status=lambda msg: console.print(msg), strict_device=strict,
|
||||
)
|
||||
|
||||
@@ -147,8 +151,10 @@ def _run_single(
|
||||
tfr = _transcribe_file(
|
||||
model=model_obj,
|
||||
actual_device=actual_device,
|
||||
backend=backend,
|
||||
model_path=model_path,
|
||||
file_path=validated_file,
|
||||
model_name=model_path,
|
||||
model_name=defaults["model"],
|
||||
compute_type=defaults["compute_type"],
|
||||
language=defaults["language"] if defaults["language"] != "auto" else None,
|
||||
on_segment=on_segment if verbose else None,
|
||||
@@ -176,12 +182,7 @@ def _run_single(
|
||||
f"Речь не обнаружена в файле {validated_file.name}", style="yellow"
|
||||
)
|
||||
|
||||
if result.device_used == "cuda":
|
||||
gpu_name = get_gpu_name()
|
||||
device_info = f"CUDA ({gpu_name or 'Unknown GPU'})"
|
||||
else:
|
||||
device_info = "CPU"
|
||||
|
||||
device_info = _format_device_info(result.device_used)
|
||||
language_mode = "detected" if defaults["language"] == "auto" else "forced"
|
||||
|
||||
content = format_transcript(
|
||||
@@ -232,15 +233,12 @@ def _run_batch(
|
||||
raise SystemExit(1)
|
||||
return
|
||||
|
||||
# Phase 2: Load model
|
||||
# Phase 2: Load model (ensure + create в одном вызове)
|
||||
requested_device = defaults["device"]
|
||||
resolved_device = detect_device(requested_device)
|
||||
strict = requested_device != "auto"
|
||||
model_path = ensure_model_available(
|
||||
defaults["model"], on_status=lambda msg: console.print(msg)
|
||||
)
|
||||
model_obj, actual_device = load_model(
|
||||
model_path, resolved_device, defaults["compute_type"],
|
||||
model_obj, actual_device, backend, model_path = load_model(
|
||||
defaults["model"], resolved_device, defaults["compute_type"],
|
||||
on_status=lambda msg: console.print(msg), strict_device=strict,
|
||||
)
|
||||
|
||||
@@ -277,8 +275,10 @@ def _run_batch(
|
||||
tfr = _transcribe_file(
|
||||
model=model_obj,
|
||||
actual_device=actual_device,
|
||||
backend=backend,
|
||||
model_path=model_path,
|
||||
file_path=file,
|
||||
model_name=model_path,
|
||||
model_name=defaults["model"],
|
||||
compute_type=defaults["compute_type"],
|
||||
language=defaults["language"] if defaults["language"] != "auto" else None,
|
||||
on_segment=on_segment if verbose else None,
|
||||
@@ -291,8 +291,11 @@ def _run_batch(
|
||||
f" {file.name}: fallback на {tfr.actual_device} при транскрипции",
|
||||
style="yellow",
|
||||
)
|
||||
# Обновляем после возможного mid-stream fallback на CPU
|
||||
model_obj, actual_device = tfr.model, tfr.actual_device
|
||||
# Обновляем после возможного mid-stream fallback
|
||||
model_obj = tfr.model
|
||||
actual_device = tfr.actual_device
|
||||
backend = tfr.backend
|
||||
model_path = tfr.model_path
|
||||
|
||||
result = tfr.result
|
||||
|
||||
@@ -301,11 +304,7 @@ def _run_batch(
|
||||
f" Речь не обнаружена: {file.name}", style="yellow"
|
||||
)
|
||||
|
||||
if result.device_used == "cuda":
|
||||
gpu_name = get_gpu_name()
|
||||
device_info = f"CUDA ({gpu_name or 'Unknown GPU'})"
|
||||
else:
|
||||
device_info = "CPU"
|
||||
device_info = _format_device_info(result.device_used)
|
||||
|
||||
content = format_transcript(
|
||||
result=result,
|
||||
|
||||
@@ -4,7 +4,7 @@ from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from .transcriber import Segment, TranscribeResult
|
||||
from .types import Segment, TranscribeResult
|
||||
|
||||
_PAUSE_THRESHOLD_S = 2.0 # пауза между сегментами для разбиения на абзацы
|
||||
_MAX_PARAGRAPH_S = 60.0 # максимальная длительность абзаца
|
||||
|
||||
@@ -1,65 +1,18 @@
|
||||
"""Обёртка над faster-whisper: загрузка моделей, транскрипция, CUDA fallback."""
|
||||
"""Оркестрация транскрипции: выбор бэкенда, загрузка модели, fallback."""
|
||||
|
||||
import warnings
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
# Должен быть ДО импорта faster_whisper / ctranslate2
|
||||
from local_transcriber._cuda_bootstrap import ensure_cublas_loadable
|
||||
from local_transcriber.backends import get_backend
|
||||
|
||||
ensure_cublas_loadable()
|
||||
|
||||
from faster_whisper import WhisperModel # noqa: E402
|
||||
from huggingface_hub import snapshot_download
|
||||
from huggingface_hub.errors import LocalEntryNotFoundError
|
||||
|
||||
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
|
||||
# Re-export из types.py для обратной совместимости
|
||||
from local_transcriber.types import ( # noqa: F401
|
||||
Segment,
|
||||
TranscribeFileResult,
|
||||
TranscribeResult,
|
||||
)
|
||||
|
||||
|
||||
def load_model(
|
||||
@@ -68,15 +21,21 @@ def load_model(
|
||||
compute_type: str,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
strict_device: bool = False,
|
||||
) -> tuple[WhisperModel, str]:
|
||||
"""Загружает модель с CUDA-фолбеком. Возвращает (model, actual_device)."""
|
||||
) -> tuple[Any, str, Any, str]:
|
||||
"""Загружает модель: ensure + create с fallback.
|
||||
|
||||
Возвращает (model, actual_device, backend, model_path).
|
||||
"""
|
||||
backend = get_backend(device)
|
||||
actual_device = device
|
||||
|
||||
model_path = backend.ensure_model_available(model_name, compute_type, on_status)
|
||||
|
||||
try:
|
||||
_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:
|
||||
# strict — пользователь явно указал устройство, fallback запрещён
|
||||
if device != "cpu" and _is_cuda_error(exc):
|
||||
if device != "cpu" and _is_backend_error(exc, device):
|
||||
if strict_device:
|
||||
raise
|
||||
warnings.warn(
|
||||
@@ -85,16 +44,21 @@ def load_model(
|
||||
stacklevel=2,
|
||||
)
|
||||
actual_device = "cpu"
|
||||
backend = get_backend("cpu")
|
||||
model_path = backend.ensure_model_available(model_name, compute_type, on_status)
|
||||
_notify_status(on_status, "Инициализирую модель на cpu...")
|
||||
model = _create_model(model_name, "cpu", compute_type)
|
||||
model = backend.create_model(model_path, "cpu", compute_type)
|
||||
else:
|
||||
raise
|
||||
return model, actual_device
|
||||
|
||||
return model, actual_device, backend, model_path
|
||||
|
||||
|
||||
def _transcribe_file(
|
||||
model: WhisperModel,
|
||||
model: Any,
|
||||
actual_device: str,
|
||||
backend: Any,
|
||||
model_path: str,
|
||||
file_path: Path,
|
||||
model_name: str,
|
||||
compute_type: str,
|
||||
@@ -103,39 +67,40 @@ def _transcribe_file(
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
strict_device: bool = False,
|
||||
) -> TranscribeFileResult:
|
||||
"""Транскрибирует один файл. При mid-stream CUDA fallback перезагружает модель."""
|
||||
"""Транскрибирует один файл. При mid-stream fallback перезагружает модель."""
|
||||
lang_arg = language if language and language != "auto" else None
|
||||
|
||||
try:
|
||||
_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:
|
||||
# Mid-stream fallback: GPU может упасть с OOM уже во время транскрипции,
|
||||
# поэтому перезагружаем модель на CPU и начинаем сначала
|
||||
if actual_device != "cpu" and _is_cuda_error(exc):
|
||||
if actual_device != "cpu" and _is_backend_error(exc, actual_device):
|
||||
if strict_device:
|
||||
raise
|
||||
warnings.warn(
|
||||
f"CUDA ошибка при транскрипции: {exc}. "
|
||||
f"Ошибка при транскрипции на {actual_device}: {exc}. "
|
||||
"Переключение на CPU и повтор.",
|
||||
stacklevel=2,
|
||||
)
|
||||
actual_device = "cpu"
|
||||
backend = get_backend("cpu")
|
||||
model_path = backend.ensure_model_available(model_name, compute_type, on_status)
|
||||
_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, "Транскрибирую...")
|
||||
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:
|
||||
raise
|
||||
|
||||
result = TranscribeResult(
|
||||
segments=segments,
|
||||
language=info.language,
|
||||
language_probability=info.language_probability,
|
||||
duration=info.duration,
|
||||
device_used=actual_device,
|
||||
return TranscribeFileResult(
|
||||
result=result,
|
||||
model=model,
|
||||
actual_device=actual_device,
|
||||
backend=backend,
|
||||
model_path=model_path,
|
||||
)
|
||||
return TranscribeFileResult(result=result, model=model, actual_device=actual_device)
|
||||
|
||||
|
||||
def transcribe(
|
||||
@@ -149,9 +114,12 @@ def transcribe(
|
||||
strict_device: bool = False,
|
||||
) -> TranscribeResult:
|
||||
"""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(
|
||||
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,
|
||||
)
|
||||
return tfr.result
|
||||
@@ -159,127 +127,30 @@ def transcribe(
|
||||
|
||||
def ensure_model_available(
|
||||
model_name: str,
|
||||
device: str = "cpu",
|
||||
compute_type: str = "float32",
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> str:
|
||||
"""Резолвит alias модели в repo_id и гарантирует наличие файлов.
|
||||
|
||||
Стратегия: cache-first (``local_files_only=True``), затем download.
|
||||
Два вызова ``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
|
||||
"""Публичный helper: гарантирует наличие модели для указанного бэкенда."""
|
||||
backend = get_backend(device)
|
||||
return backend.ensure_model_available(model_name, compute_type, on_status)
|
||||
|
||||
|
||||
def _is_cuda_error(exc: BaseException) -> bool:
|
||||
"""Проверка CUDA ошибок — используется в cli.py для Windows-диагностики."""
|
||||
msg = str(exc).lower()
|
||||
return any(k in msg for k in ("cuda", "cublas", "cudnn", "out of memory"))
|
||||
|
||||
|
||||
def _is_missing_socksio_error(exc: BaseException) -> bool:
|
||||
msg = str(exc).lower()
|
||||
return "socks proxy" in msg and "socksio" in msg
|
||||
|
||||
|
||||
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 _is_backend_error(exc: BaseException, device: str) -> bool:
|
||||
"""Определяет, связана ли ошибка с конкретным бэкендом (а не с пользовательскими данными)."""
|
||||
if device in ("cuda", "cpu"):
|
||||
return _is_cuda_error(exc)
|
||||
# openvino и другие бэкенды: конкретные паттерны ошибок добавим
|
||||
# при реализации бэкенда; пока — не маскируем ошибки
|
||||
return False
|
||||
|
||||
|
||||
def _notify_status(on_status: Callable[[str], None] | None, message: str) -> None:
|
||||
if on_status is not None:
|
||||
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
|
||||
+103
-94
@@ -24,12 +24,21 @@ def _make_model():
|
||||
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:
|
||||
result = _make_result()
|
||||
if model is None:
|
||||
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"):
|
||||
@@ -37,13 +46,13 @@ def _single_patches(result=None, tmp_file=None, actual_device="cpu"):
|
||||
if result is None:
|
||||
result = _make_result(device_used=actual_device)
|
||||
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 [
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
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.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, actual_device)),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, actual_device, backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
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")
|
||||
|
||||
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)])
|
||||
|
||||
assert out.exit_code == 0
|
||||
@@ -65,22 +74,22 @@ def test_cli_default_options_passed_to_transcribe(tmp_path):
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result()
|
||||
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)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
runner.invoke(app, [str(audio)])
|
||||
|
||||
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["language"] == "ru"
|
||||
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")
|
||||
result = _make_result(device_used="cuda")
|
||||
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)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/small")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
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]
|
||||
assert call_kwargs["model_name"] == "/models/small"
|
||||
assert call_kwargs["model_name"] == "small"
|
||||
assert call_kwargs["language"] == "ru"
|
||||
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")
|
||||
result = _make_result()
|
||||
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)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -148,7 +157,7 @@ def test_cli_empty_speech_warning(tmp_path):
|
||||
result = _make_result(segments=[])
|
||||
|
||||
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)])
|
||||
|
||||
assert out.exit_code == 0
|
||||
@@ -162,14 +171,14 @@ def test_cli_default_output_path(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript", mock_write),
|
||||
):
|
||||
@@ -187,14 +196,14 @@ def test_cli_custom_output_path(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
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")
|
||||
result = _make_result()
|
||||
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)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
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"])
|
||||
|
||||
|
||||
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.write_bytes(b"fake")
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
mock_transcribe_file = MagicMock(return_value=tfr)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/large-v3"))
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.load_model", mock_load_model),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
runner.invoke(app, [str(audio), "--model", "large-v3"])
|
||||
|
||||
mock_ensure.assert_called_once()
|
||||
call_kwargs = mock_transcribe_file.call_args[1]
|
||||
assert call_kwargs["model_name"] == "/models/large-v3"
|
||||
assert mock_load_model.call_args[0][0] == "large-v3"
|
||||
|
||||
|
||||
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.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")),
|
||||
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.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")),
|
||||
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")
|
||||
result = _make_result(device_used="cpu")
|
||||
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 (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
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")
|
||||
result = _make_result(device_used="cuda")
|
||||
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)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
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()
|
||||
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
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -366,13 +373,13 @@ def test_cli_keyboard_interrupt(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=KeyboardInterrupt),
|
||||
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.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")),
|
||||
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.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -448,14 +455,14 @@ def test_cli_batch_two_files(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
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.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -470,19 +477,18 @@ def test_cli_batch_skips_existing(tmp_path):
|
||||
b = tmp_path / "b.mp3"
|
||||
a.write_bytes(b"fake")
|
||||
b.write_bytes(b"fake")
|
||||
# Create transcript for a
|
||||
(tmp_path / "a-transcript.md").write_text("existing")
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
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.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -524,14 +530,14 @@ def test_cli_batch_force_overwrites(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
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.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -550,7 +556,8 @@ def test_cli_batch_per_file_error(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
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
|
||||
|
||||
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.validate_input_file", side_effect=lambda p: p),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=transcribe_side_effect),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -584,7 +590,8 @@ def test_cli_batch_invalid_in_prescan(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
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):
|
||||
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.validate_input_file", side_effect=validate_side_effect),
|
||||
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")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -634,42 +640,46 @@ def test_cli_config_applied(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
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 (
|
||||
patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", mock_load_model),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
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):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
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 (
|
||||
patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
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", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", mock_load_model),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
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):
|
||||
@@ -681,14 +691,14 @@ def test_cli_batch_fallback_warning(tmp_path):
|
||||
|
||||
result = _make_result(device_used="cpu")
|
||||
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 (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
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.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
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_ok = _make_result()
|
||||
model = _make_model()
|
||||
tfr_empty = _make_tfr(result=result_empty, model=model)
|
||||
tfr_ok = _make_tfr(result=result_ok, model=model)
|
||||
backend = _make_backend()
|
||||
tfr_empty = _make_tfr(result=result_empty, model=model, backend=backend)
|
||||
tfr_ok = _make_tfr(result=result_ok, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
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.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_empty, tfr_ok]),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -735,17 +745,16 @@ def test_cli_batch_midstream_fallback_warning(tmp_path):
|
||||
|
||||
model_gpu = _make_model()
|
||||
model_cpu = _make_model()
|
||||
backend = _make_backend()
|
||||
result = _make_result(device_used="cpu")
|
||||
# First file triggers mid-stream fallback
|
||||
tfr_fallback = TranscribeFileResult(result=result, model=model_cpu, actual_device="cpu")
|
||||
tfr_ok = TranscribeFileResult(result=result, model=model_cpu, actual_device="cpu")
|
||||
tfr_fallback = _make_tfr(result=result, model=model_cpu, actual_device="cpu", backend=backend)
|
||||
tfr_ok = _make_tfr(result=result, model=model_cpu, actual_device="cpu", backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
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.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_fallback, tfr_ok]),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -763,14 +772,14 @@ def test_cli_batch_model_loaded_once(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
mock_load_model = MagicMock(return_value=(model, "cpu"))
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/medium"))
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
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.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", mock_load_model),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
|
||||
+262
-293
@@ -1,9 +1,7 @@
|
||||
from collections.abc import Generator
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from huggingface_hub.errors import LocalEntryNotFoundError
|
||||
|
||||
from local_transcriber.transcriber import (
|
||||
Segment,
|
||||
@@ -15,24 +13,53 @@ from local_transcriber.transcriber import (
|
||||
)
|
||||
|
||||
|
||||
def _make_raw_segments(count: int) -> list:
|
||||
"""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
|
||||
# === Helpers ===
|
||||
|
||||
|
||||
def _make_info(language: str = "ru", probability: float = 0.95, duration: float = 60.0):
|
||||
info = MagicMock()
|
||||
info.language = language
|
||||
info.language_probability = probability
|
||||
info.duration = duration
|
||||
return info
|
||||
def _make_result(
|
||||
count: int = 2,
|
||||
language: str = "ru",
|
||||
probability: float = 0.95,
|
||||
duration: float = 60.0,
|
||||
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:
|
||||
@@ -45,14 +72,14 @@ def _create_model_dir(path: Path) -> Path:
|
||||
return path
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber.WhisperModel")
|
||||
def test_transcribe_collects_segments(mock_model_cls):
|
||||
raw_segments = _make_raw_segments(3)
|
||||
info = _make_info()
|
||||
# === transcribe() tests ===
|
||||
|
||||
instance = MagicMock()
|
||||
instance.transcribe.return_value = (iter(raw_segments), info)
|
||||
mock_model_cls.return_value = instance
|
||||
|
||||
@patch("local_transcriber.transcriber.get_backend")
|
||||
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(
|
||||
file_path=Path("test.mp3"),
|
||||
@@ -68,14 +95,11 @@ def test_transcribe_collects_segments(mock_model_cls):
|
||||
assert result.duration == 60.0
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber.WhisperModel")
|
||||
def test_transcribe_calls_on_segment(mock_model_cls):
|
||||
raw_segments = _make_raw_segments(3)
|
||||
info = _make_info()
|
||||
|
||||
instance = MagicMock()
|
||||
instance.transcribe.return_value = (iter(raw_segments), info)
|
||||
mock_model_cls.return_value = instance
|
||||
@patch("local_transcriber.transcriber.get_backend")
|
||||
def test_transcribe_calls_on_segment(mock_get_backend):
|
||||
result_data = _make_result(count=3)
|
||||
backend = _make_backend(transcribe_result=result_data)
|
||||
mock_get_backend.return_value = backend
|
||||
|
||||
callback = MagicMock()
|
||||
|
||||
@@ -86,28 +110,24 @@ def test_transcribe_calls_on_segment(mock_model_cls):
|
||||
on_segment=callback,
|
||||
)
|
||||
|
||||
assert callback.call_count == 3
|
||||
# Each call should receive a Segment instance
|
||||
for call_args in callback.call_args_list:
|
||||
seg = call_args[0][0]
|
||||
assert isinstance(seg, Segment)
|
||||
# on_segment is passed through to backend.transcribe
|
||||
call_args = backend.transcribe.call_args
|
||||
assert call_args.kwargs.get("on_segment") is callback or call_args[0][3] is callback
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber.WhisperModel")
|
||||
def test_transcribe_cuda_fallback(mock_model_cls):
|
||||
raw_segments = _make_raw_segments(2)
|
||||
info = _make_info()
|
||||
@patch("local_transcriber.transcriber.get_backend")
|
||||
def test_transcribe_cuda_fallback(mock_get_backend):
|
||||
"""CUDA error at init -> 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",
|
||||
)
|
||||
|
||||
# First call (cuda) raises, second call (cpu) succeeds
|
||||
cpu_instance = MagicMock()
|
||||
cpu_instance.transcribe.return_value = (iter(raw_segments), info)
|
||||
def backend_for_device(device):
|
||||
return cuda_backend if device == "cuda" else cpu_backend
|
||||
|
||||
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
|
||||
mock_get_backend.side_effect = backend_for_device
|
||||
|
||||
with pytest.warns(UserWarning, match="Переключение на CPU"):
|
||||
result = transcribe(
|
||||
@@ -120,14 +140,12 @@ def test_transcribe_cuda_fallback(mock_model_cls):
|
||||
assert len(result.segments) == 2
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber.WhisperModel")
|
||||
def test_transcribe_device_used(mock_model_cls):
|
||||
raw_segments = _make_raw_segments(1)
|
||||
info = _make_info()
|
||||
|
||||
instance = MagicMock()
|
||||
instance.transcribe.return_value = (iter(raw_segments), info)
|
||||
mock_model_cls.return_value = instance
|
||||
@patch("local_transcriber.transcriber.get_backend")
|
||||
def test_transcribe_device_used(mock_get_backend):
|
||||
backend = _make_backend(
|
||||
transcribe_result=_make_result(count=1, device_used="cuda"),
|
||||
)
|
||||
mock_get_backend.return_value = backend
|
||||
|
||||
result = transcribe(
|
||||
file_path=Path("test.mp3"),
|
||||
@@ -136,98 +154,47 @@ def test_transcribe_device_used(mock_model_cls):
|
||||
)
|
||||
|
||||
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")
|
||||
def test_transcribe_cuda_fallback_on_transcribe_call(mock_model_cls):
|
||||
"""CUDA error in model.transcribe() (not __init__) triggers CPU fallback."""
|
||||
raw_segments = _make_raw_segments(2)
|
||||
info = _make_info()
|
||||
|
||||
cuda_instance = MagicMock()
|
||||
cuda_instance.transcribe.side_effect = RuntimeError("CUDA error during transcription")
|
||||
|
||||
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."
|
||||
@patch("local_transcriber.transcriber.get_backend")
|
||||
def test_transcribe_cuda_fallback_on_transcribe_call(mock_get_backend):
|
||||
"""CUDA error in transcribe (not init) triggers CPU fallback."""
|
||||
cuda_backend = _make_backend(
|
||||
transcribe_error=RuntimeError("CUDA error during transcription"),
|
||||
)
|
||||
cpu_backend = _make_backend(
|
||||
transcribe_result=_make_result(count=2, device_used="cpu"),
|
||||
model_path="/mock/cpu/model",
|
||||
)
|
||||
|
||||
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(
|
||||
file_path=Path("test.mp3"),
|
||||
model_name="tiny",
|
||||
@@ -235,14 +202,10 @@ def test_transcribe_reports_missing_socksio_for_proxy(mock_model_cls):
|
||||
)
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber.WhisperModel")
|
||||
def test_transcribe_reports_status_transitions(mock_model_cls):
|
||||
raw_segments = _make_raw_segments(1)
|
||||
info = _make_info()
|
||||
|
||||
instance = MagicMock()
|
||||
instance.transcribe.return_value = (iter(raw_segments), info)
|
||||
mock_model_cls.return_value = instance
|
||||
@patch("local_transcriber.transcriber.get_backend")
|
||||
def test_transcribe_reports_status_transitions(mock_get_backend):
|
||||
backend = _make_backend(transcribe_result=_make_result(count=1))
|
||||
mock_get_backend.return_value = backend
|
||||
|
||||
statuses: list[str] = []
|
||||
|
||||
@@ -253,14 +216,138 @@ def test_transcribe_reports_status_transitions(mock_model_cls):
|
||||
on_status=statuses.append,
|
||||
)
|
||||
|
||||
assert statuses == [
|
||||
"Инициализирую модель на cpu...",
|
||||
"Транскрибирую...",
|
||||
"Транскрибирую... 00:04 / 01:00 [1 сегм.]",
|
||||
]
|
||||
# load_model reports init status, _transcribe_file reports transcribe status
|
||||
assert any("Инициализирую модель" in s for s in statuses)
|
||||
assert any("Транскрибирую" in s for s in statuses)
|
||||
|
||||
|
||||
@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):
|
||||
model_dir = _create_model_dir(tmp_path / "cache-model")
|
||||
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")
|
||||
|
||||
assert result == str(model_dir)
|
||||
mock_snapshot_download.assert_called_once_with(
|
||||
"Systran/faster-whisper-large-v3",
|
||||
local_files_only=True,
|
||||
allow_patterns=[
|
||||
"config.json",
|
||||
"preprocessor_config.json",
|
||||
"model.bin",
|
||||
"tokenizer.json",
|
||||
"vocabulary.*",
|
||||
],
|
||||
)
|
||||
mock_snapshot_download.assert_called_once()
|
||||
assert mock_snapshot_download.call_args.kwargs["local_files_only"] is True
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber._validate_model_dir")
|
||||
@patch("local_transcriber.transcriber.snapshot_download")
|
||||
@patch("local_transcriber.backends.faster_whisper._validate_model_dir")
|
||||
@patch("local_transcriber.backends.faster_whisper.snapshot_download")
|
||||
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 = [
|
||||
LocalEntryNotFoundError("not cached"),
|
||||
"/downloaded/model",
|
||||
@@ -295,10 +375,8 @@ def test_ensure_model_available_downloads_on_cache_miss(mock_snapshot_download,
|
||||
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[1].kwargs["local_files_only"] is False
|
||||
assert statuses == [
|
||||
"Проверяю кэш модели large-v3...",
|
||||
"Скачиваю модель large-v3 из Hugging Face...",
|
||||
]
|
||||
assert "Проверяю кэш модели large-v3..." in statuses
|
||||
assert "Скачиваю модель large-v3 из Hugging Face..." in statuses
|
||||
|
||||
|
||||
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):
|
||||
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")
|
||||
|
||||
assert result == str(model_dir)
|
||||
@@ -323,7 +404,7 @@ def test_ensure_model_available_rejects_unsupported_alias():
|
||||
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):
|
||||
incomplete = tmp_path / "incomplete"
|
||||
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)
|
||||
|
||||
assert result == str(complete)
|
||||
assert statuses == [
|
||||
"Проверяю кэш модели large-v3...",
|
||||
"Кэш модели large-v3 неполный, докачиваю...",
|
||||
"Скачиваю модель large-v3 из Hugging Face...",
|
||||
]
|
||||
assert "Кэш модели large-v3 неполный, докачиваю..." in statuses
|
||||
|
||||
|
||||
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="Неполная локальная модель"):
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user