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:
+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