fix(openvino): корректный compute_type в CLI + индикатор прогресса
- Зачем: - CLI показывал int8, хотя реально использовался fp16 для large-v3. - при транскрипции OpenVINO не было индикации прогресса. - Что: - _resolve_repo возвращает (repo_id, actual_compute_type). - бэкенды сохраняют actual_compute_type после ensure_model_available. - CLI выводит фактический compute_type после load_model, а не дефолтный. - generate() запускается в потоке, статус обновляется каждую секунду с elapsed time. - Проверка: - uv run pytest -q — 124 passed. - uv run transcribe file.mp4 --device openvino --model large-v3 показывает fp16. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -46,6 +46,9 @@ MODEL_REQUIRED_FILES = [
|
||||
class FasterWhisperBackend:
|
||||
"""Бэкенд транскрипции через faster-whisper (CTranslate2)."""
|
||||
|
||||
def __init__(self):
|
||||
self.actual_compute_type: str | None = None
|
||||
|
||||
def ensure_model_available(
|
||||
self,
|
||||
model_name: str,
|
||||
@@ -53,6 +56,7 @@ class FasterWhisperBackend:
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> str:
|
||||
"""Резолвит alias модели в repo_id и гарантирует наличие файлов."""
|
||||
self.actual_compute_type = compute_type
|
||||
local_path = Path(model_name).expanduser()
|
||||
if local_path.is_dir():
|
||||
_validate_model_dir(local_path)
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
import warnings
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
@@ -47,6 +49,7 @@ class OpenVINOBackend:
|
||||
def __init__(self, compute_type_explicit: bool = True):
|
||||
"""compute_type_explicit=False означает, что compute_type пришёл из дефолтов."""
|
||||
self._compute_type_explicit = compute_type_explicit
|
||||
self.actual_compute_type: str | None = None
|
||||
|
||||
def ensure_model_available(
|
||||
self,
|
||||
@@ -55,7 +58,8 @@ class OpenVINOBackend:
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> str:
|
||||
"""Скачивает/находит OpenVINO модель нужной квантизации."""
|
||||
repo_id = self._resolve_repo(model_name, compute_type)
|
||||
repo_id, resolved_ct = self._resolve_repo(model_name, compute_type)
|
||||
self.actual_compute_type = resolved_ct
|
||||
|
||||
try:
|
||||
_notify(on_status, f"Проверяю кэш модели {model_name} (OpenVINO)...")
|
||||
@@ -102,8 +106,9 @@ class OpenVINOBackend:
|
||||
if language:
|
||||
kwargs["language"] = f"<|{language}|>"
|
||||
|
||||
_notify(on_status, "Транскрибирую (OpenVINO)...")
|
||||
result = model.generate(raw_speech.tolist(), **kwargs)
|
||||
duration_str = f"{int(duration // 60):02d}:{int(duration % 60):02d}"
|
||||
pcm_list = raw_speech.tolist()
|
||||
result = _generate_with_progress(model, pcm_list, kwargs, duration_str, on_status)
|
||||
|
||||
segments: list[Segment] = []
|
||||
if hasattr(result, "chunks") and result.chunks:
|
||||
@@ -134,8 +139,11 @@ class OpenVINOBackend:
|
||||
device_used="", # оркестратор проставит
|
||||
)
|
||||
|
||||
def _resolve_repo(self, model_name: str, compute_type: str) -> str:
|
||||
"""Находит HF repo для пары (model, compute_type) с fallback."""
|
||||
def _resolve_repo(self, model_name: str, compute_type: str) -> tuple[str, str]:
|
||||
"""Находит HF repo для пары (model, compute_type) с fallback.
|
||||
|
||||
Возвращает (repo_id, actual_compute_type).
|
||||
"""
|
||||
# Для неявного compute_type: override для конкретных моделей
|
||||
if not self._compute_type_explicit and model_name in _IMPLICIT_COMPUTE_TYPE_OVERRIDES:
|
||||
compute_type = _IMPLICIT_COMPUTE_TYPE_OVERRIDES[model_name]
|
||||
@@ -143,7 +151,7 @@ class OpenVINOBackend:
|
||||
# Точное совпадение
|
||||
repo = MODEL_REPOS.get((model_name, compute_type))
|
||||
if repo:
|
||||
return repo
|
||||
return repo, compute_type
|
||||
|
||||
# Fallback только для неявного compute_type
|
||||
if not self._compute_type_explicit:
|
||||
@@ -151,7 +159,7 @@ class OpenVINOBackend:
|
||||
for fallback_ct in fallbacks:
|
||||
repo = MODEL_REPOS.get((model_name, fallback_ct))
|
||||
if repo:
|
||||
return repo
|
||||
return repo, fallback_ct
|
||||
|
||||
# Явный --compute-type с несуществующей парой → ошибка
|
||||
available = [ct for (m, ct) in MODEL_REPOS if m == model_name]
|
||||
@@ -168,6 +176,39 @@ class OpenVINOBackend:
|
||||
)
|
||||
|
||||
|
||||
def _generate_with_progress(
|
||||
model: Any,
|
||||
pcm_list: list[float],
|
||||
kwargs: dict[str, Any],
|
||||
duration_str: str,
|
||||
on_status: Callable[[str], None] | None,
|
||||
) -> Any:
|
||||
"""Запускает model.generate() в потоке, обновляя статус с elapsed time."""
|
||||
result_box: list[Any] = [None]
|
||||
error_box: list[BaseException | None] = [None]
|
||||
|
||||
def run() -> None:
|
||||
try:
|
||||
result_box[0] = model.generate(pcm_list, **kwargs)
|
||||
except BaseException as exc:
|
||||
error_box[0] = exc
|
||||
|
||||
thread = threading.Thread(target=run)
|
||||
start = time.monotonic()
|
||||
thread.start()
|
||||
|
||||
while thread.is_alive():
|
||||
elapsed = int(time.monotonic() - start)
|
||||
elapsed_str = f"{elapsed // 60:02d}:{elapsed % 60:02d}"
|
||||
_notify(on_status, f"Транскрибирую (OpenVINO)... {elapsed_str} / {duration_str} аудио")
|
||||
thread.join(timeout=1.0)
|
||||
|
||||
if error_box[0] is not None:
|
||||
raise error_box[0]
|
||||
|
||||
return result_box[0]
|
||||
|
||||
|
||||
def _notify(on_status: Callable[[str], None] | None, message: str) -> None:
|
||||
if on_status is not None:
|
||||
on_status(message)
|
||||
|
||||
@@ -136,11 +136,6 @@ def _run_single(
|
||||
output_path = build_output_path(validated_file, output)
|
||||
|
||||
console.print(f"Файл: [bold]{validated_file.name}[/bold]")
|
||||
console.print(
|
||||
f"Модель: [bold]{defaults['model']}[/bold] "
|
||||
f"Устройство: [bold]{resolved_device}[/bold] "
|
||||
f"Compute: [bold]{defaults['compute_type']}[/bold]"
|
||||
)
|
||||
|
||||
def on_segment(seg: Segment) -> None:
|
||||
console.print(f" [{seg.start:.2f}s] {seg.text.strip()}")
|
||||
@@ -150,6 +145,12 @@ def _run_single(
|
||||
on_status=lambda msg: console.print(msg), strict_device=strict,
|
||||
compute_type_explicit=compute_type_explicit,
|
||||
)
|
||||
actual_ct = getattr(backend, "actual_compute_type", defaults["compute_type"]) or defaults["compute_type"]
|
||||
console.print(
|
||||
f"Модель: [bold]{defaults['model']}[/bold] "
|
||||
f"Устройство: [bold]{actual_device}[/bold] "
|
||||
f"Compute: [bold]{actual_ct}[/bold]"
|
||||
)
|
||||
|
||||
with Status("Подготавливаю запуск...", console=console) as status:
|
||||
tfr = _transcribe_file(
|
||||
|
||||
@@ -19,12 +19,12 @@ from local_transcriber.types import Segment
|
||||
|
||||
def test_resolve_repo_exact_match():
|
||||
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||
assert backend._resolve_repo("medium", "int8") == "OpenVINO/whisper-medium-int8-ov"
|
||||
assert backend._resolve_repo("medium", "int8") == ("OpenVINO/whisper-medium-int8-ov", "int8")
|
||||
|
||||
|
||||
def test_resolve_repo_large_v3_fp16():
|
||||
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||
assert backend._resolve_repo("large-v3", "fp16") == "OpenVINO/whisper-large-v3-fp16-ov"
|
||||
assert backend._resolve_repo("large-v3", "fp16") == ("OpenVINO/whisper-large-v3-fp16-ov", "fp16")
|
||||
|
||||
|
||||
def test_resolve_repo_explicit_unsupported_pair_raises():
|
||||
@@ -44,20 +44,20 @@ 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"
|
||||
assert backend._resolve_repo("base", "int8") == ("OpenVINO/whisper-base-fp16-ov", "fp16")
|
||||
|
||||
|
||||
def test_resolve_repo_implicit_large_v3_prefers_fp16():
|
||||
"""Неявный compute_type: large-v3 автоматически получает fp16."""
|
||||
backend = OpenVINOBackend(compute_type_explicit=False)
|
||||
# Дефолт int8, но для large-v3 override на fp16
|
||||
assert backend._resolve_repo("large-v3", "int8") == "OpenVINO/whisper-large-v3-fp16-ov"
|
||||
assert backend._resolve_repo("large-v3", "int8") == ("OpenVINO/whisper-large-v3-fp16-ov", "fp16")
|
||||
|
||||
|
||||
def test_resolve_repo_explicit_large_v3_int8_respected():
|
||||
"""Явный --compute-type int8 для large-v3 → уважается."""
|
||||
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||
assert backend._resolve_repo("large-v3", "int8") == "OpenVINO/whisper-large-v3-int8-ov"
|
||||
assert backend._resolve_repo("large-v3", "int8") == ("OpenVINO/whisper-large-v3-int8-ov", "int8")
|
||||
|
||||
|
||||
# === ensure_model_available ===
|
||||
|
||||
@@ -539,5 +539,6 @@ def test_ensure_model_available_openvino_default_compute_type():
|
||||
|
||||
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||
# Проверяем что _resolve_repo работает с дефолтным compute_type для openvino (int8)
|
||||
repo = backend._resolve_repo("medium", "int8")
|
||||
repo, ct = backend._resolve_repo("medium", "int8")
|
||||
assert repo == "OpenVINO/whisper-medium-int8-ov"
|
||||
assert ct == "int8"
|
||||
|
||||
Reference in New Issue
Block a user