feat(openvino): реализован OpenVINO бэкенд транскрипции
- Зачем: - ускорение транскрипции на x86 CPU (Intel/AMD) в 2-4 раза через OpenVINO GenAI. - Что: - создан backends/openvino.py: OpenVINOBackend с ensure_model_available, create_model, transcribe. - модели скачиваются из HuggingFace (OpenVINO/whisper-*-ov), формат OpenVINO IR. - аудио декодируется через faster_whisper.decode_audio (PyAV) → .tolist() → pipe.generate(return_timestamps=True). - контракт compute_type: явный --compute-type уважается; из дефолтов large-v3 получает fp16 автоматически. - openvino-genai добавлен в pyproject.toml с platform markers (x86_64/AMD64, не macOS). - compute_type_explicit прокинут через get_backend → load_model → CLI. - 16 тестов для OpenVINO бэкенда: resolve_repo, ensure, create, transcribe, validate. - Проверка: - uv run pytest -v — 119 passed. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -10,6 +10,7 @@ dependencies = [
|
|||||||
"faster-whisper>=1.2.1",
|
"faster-whisper>=1.2.1",
|
||||||
"socksio>=1.0.0",
|
"socksio>=1.0.0",
|
||||||
"nvidia-cublas-cu12>=12.4; sys_platform == 'linux' and platform_machine == 'x86_64'",
|
"nvidia-cublas-cu12>=12.4; sys_platform == 'linux' and platform_machine == 'x86_64'",
|
||||||
|
"openvino-genai>=2025.0; sys_platform != 'darwin' and (platform_machine == 'x86_64' or platform_machine == 'AMD64')",
|
||||||
"tomli>=2.0; python_version < '3.11'",
|
"tomli>=2.0; python_version < '3.11'",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
@@ -8,10 +8,11 @@ if TYPE_CHECKING:
|
|||||||
from .base import Backend
|
from .base import Backend
|
||||||
|
|
||||||
|
|
||||||
def get_backend(device: str) -> Backend:
|
def get_backend(device: str, *, compute_type_explicit: bool = True) -> Backend:
|
||||||
"""Возвращает экземпляр бэкенда для указанного устройства.
|
"""Возвращает экземпляр бэкенда для указанного устройства.
|
||||||
|
|
||||||
Импорты ленивые — бэкенд загружается только при запросе.
|
Импорты ленивые — бэкенд загружается только при запросе.
|
||||||
|
compute_type_explicit: False если compute_type пришёл из дефолтов (влияет на fallback).
|
||||||
"""
|
"""
|
||||||
if device == "openvino":
|
if device == "openvino":
|
||||||
try:
|
try:
|
||||||
@@ -20,7 +21,7 @@ def get_backend(device: str) -> Backend:
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
"OpenVINO бэкенд недоступен. Установите: pip install openvino-genai"
|
"OpenVINO бэкенд недоступен. Установите: pip install openvino-genai"
|
||||||
) from None
|
) from None
|
||||||
return OpenVINOBackend()
|
return OpenVINOBackend(compute_type_explicit=compute_type_explicit)
|
||||||
|
|
||||||
# cuda, cpu и всё остальное → faster-whisper
|
# cuda, cpu и всё остальное → faster-whisper
|
||||||
from .faster_whisper import FasterWhisperBackend
|
from .faster_whisper import FasterWhisperBackend
|
||||||
|
|||||||
@@ -0,0 +1,179 @@
|
|||||||
|
"""Бэкенд транскрипции на основе OpenVINO GenAI."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import warnings
|
||||||
|
from collections.abc import Callable
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
from huggingface_hub.errors import LocalEntryNotFoundError
|
||||||
|
|
||||||
|
from local_transcriber.types import Segment, TranscribeResult
|
||||||
|
|
||||||
|
# (model_alias, compute_type) → HF repo
|
||||||
|
MODEL_REPOS: dict[tuple[str, str], str] = {
|
||||||
|
("tiny", "int8"): "OpenVINO/whisper-tiny-int8-ov",
|
||||||
|
("base", "fp16"): "OpenVINO/whisper-base-fp16-ov",
|
||||||
|
("small", "int8"): "OpenVINO/whisper-small-int8-ov",
|
||||||
|
("medium", "int8"): "OpenVINO/whisper-medium-int8-ov",
|
||||||
|
("large-v3", "int8"): "OpenVINO/whisper-large-v3-int8-ov",
|
||||||
|
("large-v3", "fp16"): "OpenVINO/whisper-large-v3-fp16-ov",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Fallback: если точная пара не найдена, пробуем альтернативный compute_type
|
||||||
|
_COMPUTE_TYPE_FALLBACKS: dict[str, list[str]] = {
|
||||||
|
"float32": ["fp16", "int8"],
|
||||||
|
"float16": ["fp16", "int8"],
|
||||||
|
"fp16": ["fp16", "int8"],
|
||||||
|
"int8": ["int8", "fp16"],
|
||||||
|
}
|
||||||
|
|
||||||
|
# large-v3: при неявном compute_type предпочитаем fp16 (стабильнее по качеству)
|
||||||
|
_IMPLICIT_COMPUTE_TYPE_OVERRIDES: dict[str, str] = {
|
||||||
|
"large-v3": "fp16",
|
||||||
|
}
|
||||||
|
|
||||||
|
MODEL_REQUIRED_FILES = [
|
||||||
|
"openvino_encoder_model.xml",
|
||||||
|
"openvino_decoder_model.xml",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class OpenVINOBackend:
|
||||||
|
"""Бэкенд транскрипции через openvino-genai WhisperPipeline."""
|
||||||
|
|
||||||
|
def __init__(self, compute_type_explicit: bool = True):
|
||||||
|
"""compute_type_explicit=False означает, что compute_type пришёл из дефолтов."""
|
||||||
|
self._compute_type_explicit = compute_type_explicit
|
||||||
|
|
||||||
|
def ensure_model_available(
|
||||||
|
self,
|
||||||
|
model_name: str,
|
||||||
|
compute_type: str,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
|
) -> str:
|
||||||
|
"""Скачивает/находит OpenVINO модель нужной квантизации."""
|
||||||
|
repo_id = self._resolve_repo(model_name, compute_type)
|
||||||
|
|
||||||
|
try:
|
||||||
|
_notify(on_status, f"Проверяю кэш модели {model_name} (OpenVINO)...")
|
||||||
|
cached_path = Path(snapshot_download(repo_id, local_files_only=True))
|
||||||
|
_validate_model_dir(cached_path)
|
||||||
|
return str(cached_path)
|
||||||
|
except LocalEntryNotFoundError:
|
||||||
|
pass
|
||||||
|
except ValueError:
|
||||||
|
_notify(on_status, f"Кэш модели {model_name} неполный, докачиваю...")
|
||||||
|
|
||||||
|
_notify(on_status, f"Скачиваю модель {model_name} (OpenVINO) из Hugging Face...")
|
||||||
|
downloaded_path = Path(snapshot_download(repo_id, local_files_only=False))
|
||||||
|
_validate_model_dir(downloaded_path)
|
||||||
|
return str(downloaded_path)
|
||||||
|
|
||||||
|
def create_model(
|
||||||
|
self,
|
||||||
|
model_path: str,
|
||||||
|
device: str,
|
||||||
|
compute_type: str,
|
||||||
|
) -> Any:
|
||||||
|
"""Создаёт WhisperPipeline."""
|
||||||
|
import openvino_genai as ov_genai
|
||||||
|
|
||||||
|
return ov_genai.WhisperPipeline(model_path, "CPU")
|
||||||
|
|
||||||
|
def transcribe(
|
||||||
|
self,
|
||||||
|
model: Any,
|
||||||
|
file_path: Path,
|
||||||
|
language: str | None,
|
||||||
|
on_segment: Callable[[Segment], None] | None = None,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
|
) -> TranscribeResult:
|
||||||
|
"""Транскрибирует файл через OpenVINO GenAI."""
|
||||||
|
from faster_whisper import decode_audio
|
||||||
|
|
||||||
|
_notify(on_status, "Загружаю аудио...")
|
||||||
|
raw_speech = decode_audio(str(file_path), sampling_rate=16000)
|
||||||
|
duration = len(raw_speech) / 16000.0
|
||||||
|
|
||||||
|
kwargs: dict[str, Any] = {"return_timestamps": True}
|
||||||
|
if language:
|
||||||
|
kwargs["language"] = f"<|{language}|>"
|
||||||
|
|
||||||
|
_notify(on_status, "Транскрибирую (OpenVINO)...")
|
||||||
|
result = model.generate(raw_speech.tolist(), **kwargs)
|
||||||
|
|
||||||
|
segments: list[Segment] = []
|
||||||
|
if hasattr(result, "chunks") and result.chunks:
|
||||||
|
for chunk in result.chunks:
|
||||||
|
seg = Segment(
|
||||||
|
start=chunk.start_ts,
|
||||||
|
end=chunk.end_ts,
|
||||||
|
text=chunk.text,
|
||||||
|
)
|
||||||
|
if on_segment is not None:
|
||||||
|
on_segment(seg)
|
||||||
|
segments.append(seg)
|
||||||
|
_notify(
|
||||||
|
on_status,
|
||||||
|
f"Транскрибирую (OpenVINO)... [{len(segments)} сегм.]",
|
||||||
|
)
|
||||||
|
|
||||||
|
detected_language = language or "auto"
|
||||||
|
language_probability = 1.0 if language else 0.0
|
||||||
|
|
||||||
|
return TranscribeResult(
|
||||||
|
segments=segments,
|
||||||
|
language=detected_language,
|
||||||
|
language_probability=language_probability,
|
||||||
|
duration=duration,
|
||||||
|
device_used="", # оркестратор проставит
|
||||||
|
)
|
||||||
|
|
||||||
|
def _resolve_repo(self, model_name: str, compute_type: str) -> str:
|
||||||
|
"""Находит HF repo для пары (model, compute_type) с fallback."""
|
||||||
|
# Для неявного compute_type: override для конкретных моделей
|
||||||
|
if not self._compute_type_explicit and model_name in _IMPLICIT_COMPUTE_TYPE_OVERRIDES:
|
||||||
|
compute_type = _IMPLICIT_COMPUTE_TYPE_OVERRIDES[model_name]
|
||||||
|
|
||||||
|
# Точное совпадение
|
||||||
|
repo = MODEL_REPOS.get((model_name, compute_type))
|
||||||
|
if repo:
|
||||||
|
return repo
|
||||||
|
|
||||||
|
# Fallback только для неявного compute_type
|
||||||
|
if not self._compute_type_explicit:
|
||||||
|
fallbacks = _COMPUTE_TYPE_FALLBACKS.get(compute_type, [])
|
||||||
|
for fallback_ct in fallbacks:
|
||||||
|
repo = MODEL_REPOS.get((model_name, fallback_ct))
|
||||||
|
if repo:
|
||||||
|
return repo
|
||||||
|
|
||||||
|
# Явный --compute-type с несуществующей парой → ошибка
|
||||||
|
available = [ct for (m, ct) in MODEL_REPOS if m == model_name]
|
||||||
|
if available:
|
||||||
|
raise ValueError(
|
||||||
|
f"Модель '{model_name}' недоступна с compute_type='{compute_type}' для OpenVINO. "
|
||||||
|
f"Доступные варианты: {', '.join(sorted(set(available)))}"
|
||||||
|
)
|
||||||
|
|
||||||
|
all_models = sorted({m for m, _ in MODEL_REPOS})
|
||||||
|
raise ValueError(
|
||||||
|
f"Модель '{model_name}' не найдена для OpenVINO. "
|
||||||
|
f"Доступные модели: {', '.join(all_models)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _notify(on_status: Callable[[str], None] | None, message: str) -> None:
|
||||||
|
if on_status is not None:
|
||||||
|
on_status(message)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_model_dir(model_dir: Path) -> None:
|
||||||
|
missing = [f for f in MODEL_REQUIRED_FILES if not (model_dir / f).exists()]
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
f"Неполная OpenVINO модель в '{model_dir}': отсутствуют {', '.join(missing)}"
|
||||||
|
)
|
||||||
@@ -72,6 +72,8 @@ def main(
|
|||||||
resolved_device = detect_device(defaults["device"])
|
resolved_device = detect_device(defaults["device"])
|
||||||
defaults = apply_device_defaults(defaults, resolved_device, cli_values, config)
|
defaults = apply_device_defaults(defaults, resolved_device, cli_values, config)
|
||||||
|
|
||||||
|
ct_explicit = compute_type is not None
|
||||||
|
|
||||||
expanded = expand_globs(files)
|
expanded = expand_globs(files)
|
||||||
if not expanded:
|
if not expanded:
|
||||||
console.print("Файлы не найдены.", style="red bold")
|
console.print("Файлы не найдены.", style="red bold")
|
||||||
@@ -83,9 +85,9 @@ def main(
|
|||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|
||||||
if is_batch:
|
if is_batch:
|
||||||
_run_batch(expanded, defaults, verbose, force)
|
_run_batch(expanded, defaults, verbose, force, ct_explicit)
|
||||||
else:
|
else:
|
||||||
_run_single(expanded[0], defaults, output, verbose)
|
_run_single(expanded[0], defaults, output, verbose, ct_explicit)
|
||||||
except KeyboardInterrupt:
|
except KeyboardInterrupt:
|
||||||
console.print("\nПрервано пользователем.", style="yellow")
|
console.print("\nПрервано пользователем.", style="yellow")
|
||||||
raise SystemExit(130)
|
raise SystemExit(130)
|
||||||
@@ -122,6 +124,7 @@ def _run_single(
|
|||||||
defaults: dict[str, str],
|
defaults: dict[str, str],
|
||||||
output: Path | None,
|
output: Path | None,
|
||||||
verbose: bool,
|
verbose: bool,
|
||||||
|
compute_type_explicit: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Пайплайн одного файла: валидация → модель → транскрипция → запись."""
|
"""Пайплайн одного файла: валидация → модель → транскрипция → запись."""
|
||||||
start = time.monotonic()
|
start = time.monotonic()
|
||||||
@@ -145,6 +148,7 @@ def _run_single(
|
|||||||
model_obj, actual_device, backend, model_path = load_model(
|
model_obj, actual_device, backend, model_path = load_model(
|
||||||
defaults["model"], resolved_device, defaults["compute_type"],
|
defaults["model"], resolved_device, defaults["compute_type"],
|
||||||
on_status=lambda msg: console.print(msg), strict_device=strict,
|
on_status=lambda msg: console.print(msg), strict_device=strict,
|
||||||
|
compute_type_explicit=compute_type_explicit,
|
||||||
)
|
)
|
||||||
|
|
||||||
with Status("Подготавливаю запуск...", console=console) as status:
|
with Status("Подготавливаю запуск...", console=console) as status:
|
||||||
@@ -204,6 +208,7 @@ def _run_batch(
|
|||||||
defaults: dict[str, str],
|
defaults: dict[str, str],
|
||||||
verbose: bool,
|
verbose: bool,
|
||||||
force: bool,
|
force: bool,
|
||||||
|
compute_type_explicit: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Трёхфазный батч-пайплайн: prescan → загрузка модели → транскрипция."""
|
"""Трёхфазный батч-пайплайн: prescan → загрузка модели → транскрипция."""
|
||||||
# Phase 1: Prescan — fail-fast + skip до загрузки модели (экономим ~2-5 сек)
|
# Phase 1: Prescan — fail-fast + skip до загрузки модели (экономим ~2-5 сек)
|
||||||
@@ -240,6 +245,7 @@ def _run_batch(
|
|||||||
model_obj, actual_device, backend, model_path = load_model(
|
model_obj, actual_device, backend, model_path = load_model(
|
||||||
defaults["model"], resolved_device, defaults["compute_type"],
|
defaults["model"], resolved_device, defaults["compute_type"],
|
||||||
on_status=lambda msg: console.print(msg), strict_device=strict,
|
on_status=lambda msg: console.print(msg), strict_device=strict,
|
||||||
|
compute_type_explicit=compute_type_explicit,
|
||||||
)
|
)
|
||||||
|
|
||||||
if actual_device != resolved_device:
|
if actual_device != resolved_device:
|
||||||
|
|||||||
@@ -21,12 +21,14 @@ 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,
|
||||||
|
compute_type_explicit: bool = False,
|
||||||
) -> tuple[Any, str, Any, str]:
|
) -> tuple[Any, str, Any, str]:
|
||||||
"""Загружает модель: ensure + create с fallback.
|
"""Загружает модель: ensure + create с fallback.
|
||||||
|
|
||||||
Возвращает (model, actual_device, backend, model_path).
|
Возвращает (model, actual_device, backend, model_path).
|
||||||
|
compute_type_explicit: True если пользователь явно указал --compute-type.
|
||||||
"""
|
"""
|
||||||
backend = get_backend(device)
|
backend = get_backend(device, compute_type_explicit=compute_type_explicit)
|
||||||
actual_device = device
|
actual_device = device
|
||||||
|
|
||||||
model_path = backend.ensure_model_available(model_name, compute_type, on_status)
|
model_path = backend.ensure_model_available(model_name, compute_type, on_status)
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
"""Тесты для OpenVINO бэкенда."""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from local_transcriber.backends.openvino import (
|
||||||
|
MODEL_REPOS,
|
||||||
|
OpenVINOBackend,
|
||||||
|
_validate_model_dir,
|
||||||
|
)
|
||||||
|
from local_transcriber.types import Segment
|
||||||
|
|
||||||
|
|
||||||
|
# === _resolve_repo ===
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_repo_exact_match():
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||||
|
assert backend._resolve_repo("medium", "int8") == "OpenVINO/whisper-medium-int8-ov"
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_repo_large_v3_fp16():
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||||
|
assert backend._resolve_repo("large-v3", "fp16") == "OpenVINO/whisper-large-v3-fp16-ov"
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_repo_explicit_unsupported_pair_raises():
|
||||||
|
"""Явный --compute-type с несуществующей парой → ошибка."""
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||||
|
with pytest.raises(ValueError, match="недоступна с compute_type='fp16'"):
|
||||||
|
backend._resolve_repo("medium", "fp16")
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_repo_explicit_unknown_model_raises():
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||||
|
with pytest.raises(ValueError, match="не найдена для OpenVINO"):
|
||||||
|
backend._resolve_repo("distil-large-v3", "int8")
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_repo_implicit_fallback():
|
||||||
|
"""Неявный compute_type: если int8 недоступен для base, fallback на fp16."""
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=False)
|
||||||
|
# base + int8 не существует, но base + fp16 есть
|
||||||
|
assert backend._resolve_repo("base", "int8") == "OpenVINO/whisper-base-fp16-ov"
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_repo_implicit_large_v3_prefers_fp16():
|
||||||
|
"""Неявный compute_type: large-v3 автоматически получает fp16."""
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=False)
|
||||||
|
# Дефолт int8, но для large-v3 override на fp16
|
||||||
|
assert backend._resolve_repo("large-v3", "int8") == "OpenVINO/whisper-large-v3-fp16-ov"
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_repo_explicit_large_v3_int8_respected():
|
||||||
|
"""Явный --compute-type int8 для large-v3 → уважается."""
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||||
|
assert backend._resolve_repo("large-v3", "int8") == "OpenVINO/whisper-large-v3-int8-ov"
|
||||||
|
|
||||||
|
|
||||||
|
# === ensure_model_available ===
|
||||||
|
|
||||||
|
|
||||||
|
@patch("local_transcriber.backends.openvino.snapshot_download")
|
||||||
|
def test_ensure_model_available_cache_hit(mock_download, tmp_path):
|
||||||
|
model_dir = tmp_path / "model"
|
||||||
|
model_dir.mkdir()
|
||||||
|
(model_dir / "openvino_encoder_model.xml").write_text("<xml/>")
|
||||||
|
(model_dir / "openvino_decoder_model.xml").write_text("<xml/>")
|
||||||
|
mock_download.return_value = str(model_dir)
|
||||||
|
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||||
|
result = backend.ensure_model_available("medium", "int8")
|
||||||
|
|
||||||
|
assert result == str(model_dir)
|
||||||
|
mock_download.assert_called_once()
|
||||||
|
assert mock_download.call_args.kwargs["local_files_only"] is True
|
||||||
|
|
||||||
|
|
||||||
|
@patch("local_transcriber.backends.openvino.snapshot_download")
|
||||||
|
def test_ensure_model_available_downloads(mock_download, tmp_path):
|
||||||
|
from huggingface_hub.errors import LocalEntryNotFoundError
|
||||||
|
|
||||||
|
model_dir = tmp_path / "downloaded"
|
||||||
|
model_dir.mkdir()
|
||||||
|
(model_dir / "openvino_encoder_model.xml").write_text("<xml/>")
|
||||||
|
(model_dir / "openvino_decoder_model.xml").write_text("<xml/>")
|
||||||
|
|
||||||
|
mock_download.side_effect = [
|
||||||
|
LocalEntryNotFoundError("not cached"),
|
||||||
|
str(model_dir),
|
||||||
|
]
|
||||||
|
|
||||||
|
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||||
|
statuses: list[str] = []
|
||||||
|
result = backend.ensure_model_available("medium", "int8", on_status=statuses.append)
|
||||||
|
|
||||||
|
assert result == str(model_dir)
|
||||||
|
assert any("Скачиваю" in s for s in statuses)
|
||||||
|
|
||||||
|
|
||||||
|
# === create_model ===
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_model():
|
||||||
|
mock_ov = MagicMock()
|
||||||
|
mock_pipeline = MagicMock()
|
||||||
|
mock_ov.WhisperPipeline.return_value = mock_pipeline
|
||||||
|
|
||||||
|
backend = OpenVINOBackend()
|
||||||
|
with patch.dict("sys.modules", {"openvino_genai": mock_ov}):
|
||||||
|
model = backend.create_model("/path/to/model", "openvino", "int8")
|
||||||
|
|
||||||
|
mock_ov.WhisperPipeline.assert_called_once_with("/path/to/model", "CPU")
|
||||||
|
assert model is mock_pipeline
|
||||||
|
|
||||||
|
|
||||||
|
# === transcribe ===
|
||||||
|
|
||||||
|
|
||||||
|
def test_transcribe_maps_chunks_to_segments():
|
||||||
|
"""Проверяет маппинг chunks → Segment[] и формат языка."""
|
||||||
|
backend = OpenVINOBackend()
|
||||||
|
|
||||||
|
mock_model = MagicMock()
|
||||||
|
chunk1 = MagicMock()
|
||||||
|
chunk1.start_ts = 0.0
|
||||||
|
chunk1.end_ts = 3.5
|
||||||
|
chunk1.text = " Привет мир"
|
||||||
|
chunk2 = MagicMock()
|
||||||
|
chunk2.start_ts = 3.5
|
||||||
|
chunk2.end_ts = 7.0
|
||||||
|
chunk2.text = " Тестовый сегмент"
|
||||||
|
|
||||||
|
mock_result = MagicMock()
|
||||||
|
mock_result.chunks = [chunk1, chunk2]
|
||||||
|
mock_model.generate.return_value = mock_result
|
||||||
|
|
||||||
|
raw_audio = np.zeros(16000 * 10, dtype=np.float32) # 10 секунд
|
||||||
|
|
||||||
|
with patch("faster_whisper.decode_audio", return_value=raw_audio):
|
||||||
|
result = backend.transcribe(
|
||||||
|
mock_model, Path("test.mp3"), language="ru",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(result.segments) == 2
|
||||||
|
assert result.segments[0].text == " Привет мир"
|
||||||
|
assert result.segments[0].start == 0.0
|
||||||
|
assert result.segments[0].end == 3.5
|
||||||
|
assert result.duration == 10.0
|
||||||
|
|
||||||
|
# Проверяем формат языка для OpenVINO GenAI
|
||||||
|
call_kwargs = mock_model.generate.call_args
|
||||||
|
assert call_kwargs.kwargs["language"] == "<|ru|>"
|
||||||
|
assert call_kwargs.kwargs["return_timestamps"] is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_transcribe_calls_tolist():
|
||||||
|
"""raw_speech передаётся как list, не ndarray."""
|
||||||
|
backend = OpenVINOBackend()
|
||||||
|
mock_model = MagicMock()
|
||||||
|
mock_result = MagicMock()
|
||||||
|
mock_result.chunks = []
|
||||||
|
mock_model.generate.return_value = mock_result
|
||||||
|
|
||||||
|
raw_audio = np.zeros(160, dtype=np.float32)
|
||||||
|
|
||||||
|
with patch("faster_whisper.decode_audio", return_value=raw_audio):
|
||||||
|
backend.transcribe(mock_model, Path("test.mp3"), language=None)
|
||||||
|
|
||||||
|
call_args = mock_model.generate.call_args[0][0]
|
||||||
|
assert isinstance(call_args, list)
|
||||||
|
|
||||||
|
|
||||||
|
def test_transcribe_no_language_auto():
|
||||||
|
"""Без указания языка — не передаём language в generate."""
|
||||||
|
backend = OpenVINOBackend()
|
||||||
|
mock_model = MagicMock()
|
||||||
|
mock_result = MagicMock()
|
||||||
|
mock_result.chunks = []
|
||||||
|
mock_model.generate.return_value = mock_result
|
||||||
|
|
||||||
|
raw_audio = np.zeros(160, dtype=np.float32)
|
||||||
|
|
||||||
|
with patch("faster_whisper.decode_audio", return_value=raw_audio):
|
||||||
|
result = backend.transcribe(mock_model, Path("test.mp3"), language=None)
|
||||||
|
|
||||||
|
call_kwargs = mock_model.generate.call_args.kwargs
|
||||||
|
assert "language" not in call_kwargs
|
||||||
|
assert result.language == "auto"
|
||||||
|
assert result.language_probability == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
def test_transcribe_calls_on_segment():
|
||||||
|
backend = OpenVINOBackend()
|
||||||
|
mock_model = MagicMock()
|
||||||
|
chunk = MagicMock()
|
||||||
|
chunk.start_ts = 0.0
|
||||||
|
chunk.end_ts = 2.0
|
||||||
|
chunk.text = " Test"
|
||||||
|
mock_result = MagicMock()
|
||||||
|
mock_result.chunks = [chunk]
|
||||||
|
mock_model.generate.return_value = mock_result
|
||||||
|
|
||||||
|
raw_audio = np.zeros(16000, dtype=np.float32)
|
||||||
|
callback = MagicMock()
|
||||||
|
|
||||||
|
with patch("faster_whisper.decode_audio", return_value=raw_audio):
|
||||||
|
backend.transcribe(
|
||||||
|
mock_model, Path("test.mp3"), language="en", on_segment=callback,
|
||||||
|
)
|
||||||
|
|
||||||
|
callback.assert_called_once()
|
||||||
|
seg = callback.call_args[0][0]
|
||||||
|
assert isinstance(seg, Segment)
|
||||||
|
assert seg.text == " Test"
|
||||||
|
|
||||||
|
|
||||||
|
# === _validate_model_dir ===
|
||||||
|
|
||||||
|
|
||||||
|
def test_validate_model_dir_ok(tmp_path):
|
||||||
|
(tmp_path / "openvino_encoder_model.xml").write_text("<xml/>")
|
||||||
|
(tmp_path / "openvino_decoder_model.xml").write_text("<xml/>")
|
||||||
|
_validate_model_dir(tmp_path) # should not raise
|
||||||
|
|
||||||
|
|
||||||
|
def test_validate_model_dir_missing(tmp_path):
|
||||||
|
(tmp_path / "openvino_encoder_model.xml").write_text("<xml/>")
|
||||||
|
with pytest.raises(ValueError, match="openvino_decoder_model.xml"):
|
||||||
|
_validate_model_dir(tmp_path)
|
||||||
@@ -124,7 +124,7 @@ def test_transcribe_cuda_fallback(mock_get_backend):
|
|||||||
model_path="/mock/cpu/model",
|
model_path="/mock/cpu/model",
|
||||||
)
|
)
|
||||||
|
|
||||||
def backend_for_device(device):
|
def backend_for_device(device, **kwargs):
|
||||||
return cuda_backend if device == "cuda" else cpu_backend
|
return cuda_backend if device == "cuda" else cpu_backend
|
||||||
|
|
||||||
mock_get_backend.side_effect = backend_for_device
|
mock_get_backend.side_effect = backend_for_device
|
||||||
@@ -168,7 +168,7 @@ def test_transcribe_cuda_fallback_on_transcribe_call(mock_get_backend):
|
|||||||
model_path="/mock/cpu/model",
|
model_path="/mock/cpu/model",
|
||||||
)
|
)
|
||||||
|
|
||||||
def backend_for_device(device):
|
def backend_for_device(device, **kwargs):
|
||||||
return cuda_backend if device == "cuda" else cpu_backend
|
return cuda_backend if device == "cuda" else cpu_backend
|
||||||
|
|
||||||
mock_get_backend.side_effect = backend_for_device
|
mock_get_backend.side_effect = backend_for_device
|
||||||
@@ -245,7 +245,7 @@ def test_transcribe_non_strict_cuda_fallback(mock_get_backend):
|
|||||||
model_path="/mock/cpu/model",
|
model_path="/mock/cpu/model",
|
||||||
)
|
)
|
||||||
|
|
||||||
def backend_for_device(device):
|
def backend_for_device(device, **kwargs):
|
||||||
return cuda_backend if device == "cuda" else cpu_backend
|
return cuda_backend if device == "cuda" else cpu_backend
|
||||||
|
|
||||||
mock_get_backend.side_effect = backend_for_device
|
mock_get_backend.side_effect = backend_for_device
|
||||||
@@ -288,7 +288,7 @@ def test_load_model_cuda_fallback(mock_get_backend):
|
|||||||
cpu_model = MagicMock()
|
cpu_model = MagicMock()
|
||||||
cpu_backend = _make_backend(model=cpu_model, model_path="/mock/cpu/model")
|
cpu_backend = _make_backend(model=cpu_model, model_path="/mock/cpu/model")
|
||||||
|
|
||||||
def backend_for_device(device):
|
def backend_for_device(device, **kwargs):
|
||||||
return cuda_backend if device == "cuda" else cpu_backend
|
return cuda_backend if device == "cuda" else cpu_backend
|
||||||
|
|
||||||
mock_get_backend.side_effect = backend_for_device
|
mock_get_backend.side_effect = backend_for_device
|
||||||
|
|||||||
Reference in New Issue
Block a user