feat(onnx): добавлен каталог моделей GigaAM
- Зачем: - добавлена локальная транскрипция смешанной речи и русского текста с пунктуацией. - Что: - onnx-asr обновлён до 0.12 и зарегистрированы три модели GigaAM. - выбор compute_type учитывает опубликованные квантизации и явность настройки. - обновлены тесты, README и требования PRD. - Проверка: - uv run pytest -q: 223 passed, 1 skipped. - выполнены smoke- и полные прогоны четырёх GigaAM-моделей.
This commit is contained in:
@@ -32,7 +32,7 @@ def get_backend(device: str, *, compute_type_explicit: bool = True) -> Backend:
|
||||
raise ValueError(
|
||||
"onnx-asr бэкенд недоступен. Установите: pip install onnx-asr[cpu,hub]"
|
||||
) from None
|
||||
return OnnxAsrBackend()
|
||||
return OnnxAsrBackend(compute_type_explicit=compute_type_explicit)
|
||||
|
||||
# cuda, cpu и всё остальное → faster-whisper
|
||||
from .faster_whisper import FasterWhisperBackend
|
||||
|
||||
@@ -3,14 +3,46 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from local_transcriber.types import Segment, TranscribeResult
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class OnnxModelSpec:
|
||||
"""Имя onnx-asr и опубликованные варианты квантизации модели."""
|
||||
|
||||
model_id: str
|
||||
quantizations: frozenset[str | None]
|
||||
|
||||
|
||||
_INT8_AND_FLOAT32 = frozenset({"int8", None})
|
||||
|
||||
MODEL_CATALOG: dict[str, OnnxModelSpec] = {
|
||||
"gigaam-v3": OnnxModelSpec("gigaam-v3-ctc", _INT8_AND_FLOAT32),
|
||||
"parakeet-v3": OnnxModelSpec(
|
||||
"nemo-parakeet-tdt-0.6b-v3",
|
||||
_INT8_AND_FLOAT32,
|
||||
),
|
||||
"gigaam-multilingual-ctc": OnnxModelSpec(
|
||||
"gigaam-multilingual-ctc",
|
||||
_INT8_AND_FLOAT32,
|
||||
),
|
||||
"gigaam-v3-e2e-ctc": OnnxModelSpec(
|
||||
"gigaam-v3-e2e-ctc",
|
||||
_INT8_AND_FLOAT32,
|
||||
),
|
||||
"gigaam-v3-e2e-rnnt": OnnxModelSpec(
|
||||
"gigaam-v3-e2e-rnnt",
|
||||
_INT8_AND_FLOAT32,
|
||||
),
|
||||
}
|
||||
|
||||
# Оставлено как совместимое представление публичного каталога алиасов.
|
||||
MODEL_ALIASES: dict[str, str] = {
|
||||
"gigaam-v3": "gigaam-v3-ctc",
|
||||
"parakeet-v3": "nemo-parakeet-tdt-0.6b-v3",
|
||||
alias: spec.model_id for alias, spec in MODEL_CATALOG.items()
|
||||
}
|
||||
|
||||
SUPPORTED_ALIASES = ", ".join(MODEL_ALIASES)
|
||||
@@ -45,7 +77,8 @@ def _normalize_quantization(compute_type: str) -> str | None:
|
||||
class OnnxAsrBackend:
|
||||
"""Бэкенд транскрипции через onnx-asr (ONNX Runtime)."""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, compute_type_explicit: bool = True):
|
||||
self._compute_type_explicit = compute_type_explicit
|
||||
self.actual_compute_type: str | None = None
|
||||
self._resolved_model_id: str | None = None
|
||||
self._vad: Any = None
|
||||
@@ -61,7 +94,26 @@ class OnnxAsrBackend:
|
||||
onnx-asr downloads models automatically via load_model(),
|
||||
so this just validates the alias and returns the identifier string.
|
||||
"""
|
||||
self.actual_compute_type = compute_type
|
||||
spec = MODEL_CATALOG.get(model_name)
|
||||
quantization = _normalize_quantization(compute_type)
|
||||
if spec is not None and quantization not in spec.quantizations:
|
||||
if self._compute_type_explicit:
|
||||
available = _format_compute_types(spec.quantizations)
|
||||
raise ValueError(
|
||||
f"Модель '{model_name}' недоступна с compute_type='{compute_type}' "
|
||||
f"для onnx-asr. Доступные варианты: {available}"
|
||||
)
|
||||
|
||||
resolved_compute_type = _preferred_compute_type(spec.quantizations)
|
||||
_notify(
|
||||
on_status,
|
||||
f"Модель {model_name} недоступна с compute_type={compute_type}; "
|
||||
f"использую {resolved_compute_type}.",
|
||||
)
|
||||
else:
|
||||
resolved_compute_type = _compute_type_for_quantization(quantization)
|
||||
|
||||
self.actual_compute_type = resolved_compute_type
|
||||
self._resolved_model_id = self._resolve_model(model_name)
|
||||
return self._resolved_model_id
|
||||
|
||||
@@ -79,7 +131,8 @@ class OnnxAsrBackend:
|
||||
"""
|
||||
import onnx_asr
|
||||
|
||||
quantization = _normalize_quantization(compute_type)
|
||||
actual_compute_type = self.actual_compute_type or compute_type
|
||||
quantization = _normalize_quantization(actual_compute_type)
|
||||
|
||||
model = onnx_asr.load_model(
|
||||
model=model_path,
|
||||
@@ -113,7 +166,9 @@ class OnnxAsrBackend:
|
||||
segments: list[Segment] = []
|
||||
detected_language = language or "unknown"
|
||||
|
||||
for vad_seg in model.recognize(audio_array, sample_rate=16000, language=language):
|
||||
for vad_seg in model.recognize(
|
||||
audio_array, sample_rate=16000, language=language
|
||||
):
|
||||
seg = Segment(
|
||||
start=max(0.0, vad_seg.start),
|
||||
end=max(0.0, vad_seg.end),
|
||||
@@ -149,6 +204,22 @@ class OnnxAsrBackend:
|
||||
)
|
||||
|
||||
|
||||
def _format_compute_types(quantizations: frozenset[str | None]) -> str:
|
||||
values = [_compute_type_for_quantization(value) for value in quantizations]
|
||||
return ", ".join(sorted(values))
|
||||
|
||||
|
||||
def _compute_type_for_quantization(quantization: str | None) -> str:
|
||||
return "float32" if quantization is None else quantization
|
||||
|
||||
|
||||
def _preferred_compute_type(quantizations: frozenset[str | None]) -> str:
|
||||
for quantization in ("int8", None, "fp16"):
|
||||
if quantization in quantizations:
|
||||
return _compute_type_for_quantization(quantization)
|
||||
raise ValueError("Для ONNX-модели не указаны доступные квантизации")
|
||||
|
||||
|
||||
def _notify(on_status: Callable[[str], None] | None, message: str) -> None:
|
||||
if on_status is not None:
|
||||
on_status(message)
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
"""Загрузка конфигурации из ``.transcriber.toml`` и каскад приоритетов."""
|
||||
|
||||
import tomllib
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
|
||||
import tomllib
|
||||
|
||||
HARDCODED_DEFAULTS: dict[str, str] = {
|
||||
"model": "medium",
|
||||
"language": "ru",
|
||||
@@ -23,7 +22,15 @@ DEVICE_DEFAULTS: dict[str, dict[str, str]] = {
|
||||
|
||||
# Одно место правды для допустимых ключей конфига
|
||||
_VALID_KEYS = set(HARDCODED_DEFAULTS)
|
||||
_VALID_DEVICES = {"auto", "cpu", "cuda", "openvino", "openvino-gpu", "openvino-cpu", "onnx"}
|
||||
_VALID_DEVICES = {
|
||||
"auto",
|
||||
"cpu",
|
||||
"cuda",
|
||||
"openvino",
|
||||
"openvino-gpu",
|
||||
"openvino-cpu",
|
||||
"onnx",
|
||||
}
|
||||
|
||||
|
||||
def find_config_file() -> Path | None:
|
||||
|
||||
Reference in New Issue
Block a user