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:
|
class FasterWhisperBackend:
|
||||||
"""Бэкенд транскрипции через faster-whisper (CTranslate2)."""
|
"""Бэкенд транскрипции через faster-whisper (CTranslate2)."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.actual_compute_type: str | None = None
|
||||||
|
|
||||||
def ensure_model_available(
|
def ensure_model_available(
|
||||||
self,
|
self,
|
||||||
model_name: str,
|
model_name: str,
|
||||||
@@ -53,6 +56,7 @@ class FasterWhisperBackend:
|
|||||||
on_status: Callable[[str], None] | None = None,
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Резолвит alias модели в repo_id и гарантирует наличие файлов."""
|
"""Резолвит alias модели в repo_id и гарантирует наличие файлов."""
|
||||||
|
self.actual_compute_type = compute_type
|
||||||
local_path = Path(model_name).expanduser()
|
local_path = Path(model_name).expanduser()
|
||||||
if local_path.is_dir():
|
if local_path.is_dir():
|
||||||
_validate_model_dir(local_path)
|
_validate_model_dir(local_path)
|
||||||
|
|||||||
@@ -2,6 +2,8 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
import warnings
|
import warnings
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -47,6 +49,7 @@ class OpenVINOBackend:
|
|||||||
def __init__(self, compute_type_explicit: bool = True):
|
def __init__(self, compute_type_explicit: bool = True):
|
||||||
"""compute_type_explicit=False означает, что compute_type пришёл из дефолтов."""
|
"""compute_type_explicit=False означает, что compute_type пришёл из дефолтов."""
|
||||||
self._compute_type_explicit = compute_type_explicit
|
self._compute_type_explicit = compute_type_explicit
|
||||||
|
self.actual_compute_type: str | None = None
|
||||||
|
|
||||||
def ensure_model_available(
|
def ensure_model_available(
|
||||||
self,
|
self,
|
||||||
@@ -55,7 +58,8 @@ class OpenVINOBackend:
|
|||||||
on_status: Callable[[str], None] | None = None,
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Скачивает/находит OpenVINO модель нужной квантизации."""
|
"""Скачивает/находит 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:
|
try:
|
||||||
_notify(on_status, f"Проверяю кэш модели {model_name} (OpenVINO)...")
|
_notify(on_status, f"Проверяю кэш модели {model_name} (OpenVINO)...")
|
||||||
@@ -102,8 +106,9 @@ class OpenVINOBackend:
|
|||||||
if language:
|
if language:
|
||||||
kwargs["language"] = f"<|{language}|>"
|
kwargs["language"] = f"<|{language}|>"
|
||||||
|
|
||||||
_notify(on_status, "Транскрибирую (OpenVINO)...")
|
duration_str = f"{int(duration // 60):02d}:{int(duration % 60):02d}"
|
||||||
result = model.generate(raw_speech.tolist(), **kwargs)
|
pcm_list = raw_speech.tolist()
|
||||||
|
result = _generate_with_progress(model, pcm_list, kwargs, duration_str, on_status)
|
||||||
|
|
||||||
segments: list[Segment] = []
|
segments: list[Segment] = []
|
||||||
if hasattr(result, "chunks") and result.chunks:
|
if hasattr(result, "chunks") and result.chunks:
|
||||||
@@ -134,8 +139,11 @@ class OpenVINOBackend:
|
|||||||
device_used="", # оркестратор проставит
|
device_used="", # оркестратор проставит
|
||||||
)
|
)
|
||||||
|
|
||||||
def _resolve_repo(self, model_name: str, compute_type: str) -> str:
|
def _resolve_repo(self, model_name: str, compute_type: str) -> tuple[str, str]:
|
||||||
"""Находит HF repo для пары (model, compute_type) с fallback."""
|
"""Находит HF repo для пары (model, compute_type) с fallback.
|
||||||
|
|
||||||
|
Возвращает (repo_id, actual_compute_type).
|
||||||
|
"""
|
||||||
# Для неявного compute_type: override для конкретных моделей
|
# Для неявного compute_type: override для конкретных моделей
|
||||||
if not self._compute_type_explicit and model_name in _IMPLICIT_COMPUTE_TYPE_OVERRIDES:
|
if not self._compute_type_explicit and model_name in _IMPLICIT_COMPUTE_TYPE_OVERRIDES:
|
||||||
compute_type = _IMPLICIT_COMPUTE_TYPE_OVERRIDES[model_name]
|
compute_type = _IMPLICIT_COMPUTE_TYPE_OVERRIDES[model_name]
|
||||||
@@ -143,7 +151,7 @@ class OpenVINOBackend:
|
|||||||
# Точное совпадение
|
# Точное совпадение
|
||||||
repo = MODEL_REPOS.get((model_name, compute_type))
|
repo = MODEL_REPOS.get((model_name, compute_type))
|
||||||
if repo:
|
if repo:
|
||||||
return repo
|
return repo, compute_type
|
||||||
|
|
||||||
# Fallback только для неявного compute_type
|
# Fallback только для неявного compute_type
|
||||||
if not self._compute_type_explicit:
|
if not self._compute_type_explicit:
|
||||||
@@ -151,7 +159,7 @@ class OpenVINOBackend:
|
|||||||
for fallback_ct in fallbacks:
|
for fallback_ct in fallbacks:
|
||||||
repo = MODEL_REPOS.get((model_name, fallback_ct))
|
repo = MODEL_REPOS.get((model_name, fallback_ct))
|
||||||
if repo:
|
if repo:
|
||||||
return repo
|
return repo, fallback_ct
|
||||||
|
|
||||||
# Явный --compute-type с несуществующей парой → ошибка
|
# Явный --compute-type с несуществующей парой → ошибка
|
||||||
available = [ct for (m, ct) in MODEL_REPOS if m == model_name]
|
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:
|
def _notify(on_status: Callable[[str], None] | None, message: str) -> None:
|
||||||
if on_status is not None:
|
if on_status is not None:
|
||||||
on_status(message)
|
on_status(message)
|
||||||
|
|||||||
@@ -136,11 +136,6 @@ def _run_single(
|
|||||||
output_path = build_output_path(validated_file, output)
|
output_path = build_output_path(validated_file, output)
|
||||||
|
|
||||||
console.print(f"Файл: [bold]{validated_file.name}[/bold]")
|
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:
|
def on_segment(seg: Segment) -> None:
|
||||||
console.print(f" [{seg.start:.2f}s] {seg.text.strip()}")
|
console.print(f" [{seg.start:.2f}s] {seg.text.strip()}")
|
||||||
@@ -150,6 +145,12 @@ def _run_single(
|
|||||||
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,
|
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:
|
with Status("Подготавливаю запуск...", console=console) as status:
|
||||||
tfr = _transcribe_file(
|
tfr = _transcribe_file(
|
||||||
|
|||||||
@@ -19,12 +19,12 @@ from local_transcriber.types import Segment
|
|||||||
|
|
||||||
def test_resolve_repo_exact_match():
|
def test_resolve_repo_exact_match():
|
||||||
backend = OpenVINOBackend(compute_type_explicit=True)
|
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():
|
def test_resolve_repo_large_v3_fp16():
|
||||||
backend = OpenVINOBackend(compute_type_explicit=True)
|
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():
|
def test_resolve_repo_explicit_unsupported_pair_raises():
|
||||||
@@ -44,20 +44,20 @@ def test_resolve_repo_implicit_fallback():
|
|||||||
"""Неявный compute_type: если int8 недоступен для base, fallback на fp16."""
|
"""Неявный compute_type: если int8 недоступен для base, fallback на fp16."""
|
||||||
backend = OpenVINOBackend(compute_type_explicit=False)
|
backend = OpenVINOBackend(compute_type_explicit=False)
|
||||||
# base + int8 не существует, но base + fp16 есть
|
# 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():
|
def test_resolve_repo_implicit_large_v3_prefers_fp16():
|
||||||
"""Неявный compute_type: large-v3 автоматически получает fp16."""
|
"""Неявный compute_type: large-v3 автоматически получает fp16."""
|
||||||
backend = OpenVINOBackend(compute_type_explicit=False)
|
backend = OpenVINOBackend(compute_type_explicit=False)
|
||||||
# Дефолт int8, но для large-v3 override на fp16
|
# Дефолт 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():
|
def test_resolve_repo_explicit_large_v3_int8_respected():
|
||||||
"""Явный --compute-type int8 для large-v3 → уважается."""
|
"""Явный --compute-type int8 для large-v3 → уважается."""
|
||||||
backend = OpenVINOBackend(compute_type_explicit=True)
|
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 ===
|
# === ensure_model_available ===
|
||||||
|
|||||||
@@ -539,5 +539,6 @@ def test_ensure_model_available_openvino_default_compute_type():
|
|||||||
|
|
||||||
backend = OpenVINOBackend(compute_type_explicit=True)
|
backend = OpenVINOBackend(compute_type_explicit=True)
|
||||||
# Проверяем что _resolve_repo работает с дефолтным compute_type для openvino (int8)
|
# Проверяем что _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 repo == "OpenVINO/whisper-medium-int8-ov"
|
||||||
|
assert ct == "int8"
|
||||||
|
|||||||
Reference in New Issue
Block a user