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:
Dmitriy Dementiev
2026-08-11 13:11:18 +03:00
parent 3c519c6860
commit 0af2dbdd17
8 changed files with 177 additions and 27 deletions
+1 -1
View File
@@ -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
+77 -6
View File
@@ -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)
+10 -3
View File
@@ -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: