Files
local-transcriber/src/local_transcriber/backends/onnx_asr.py
T
Dmitriy Dementiev 8f85930616 feat(onnx): добавлена модель GigaAM Multilingual Large
- Зачем:
  - расширенный benchmark показал устойчивое улучшение multilingual large на трёх реальных записях.
- Что:
  - добавлен alias gigaam-multilingual-large-ctc с квантизациями int8 и float32.
  - обновлены README, спецификация, backlog и сравнительный benchmark.
- Проверка:
  - uv run pytest -q: 225 passed, 1 skipped.
  - uvx ruff check src/local_transcriber/backends/onnx_asr.py tests/test_onnx_asr.py.
2026-08-11 16:12:36 +03:00

234 lines
8.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Бэкенд транскрипции на основе onnx-asr (GigaAM, Parakeet, FastConformer)."""
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-multilingual-large-ctc": OnnxModelSpec(
"gigaam-multilingual-large-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] = {
alias: spec.model_id for alias, spec in MODEL_CATALOG.items()
}
SUPPORTED_ALIASES = ", ".join(MODEL_ALIASES)
# compute_type проекта → onnx-asr quantization (file suffix; None = unquantized).
_QUANTIZATION_MAP: dict[str, str | None] = {
"int8": "int8",
"fp16": "fp16",
"float16": "fp16",
"float32": None,
"fp32": None,
}
def _normalize_quantization(compute_type: str) -> str | None:
"""Маппит compute_type проекта в значение onnx-asr ``quantization``.
onnx-asr использует ``quantization`` как суффикс имени файла модели:
``int8``/``fp16`` подгружают квантизованные веса, ``None`` — unquantized
(float32). Передача ``"float32"`` строкой пытается найти несуществующий
файл с суффиксом ``_float32`` и приводит к ошибке загрузки.
"""
if compute_type not in _QUANTIZATION_MAP:
supported = ", ".join(sorted(_QUANTIZATION_MAP))
raise ValueError(
f"Неподдерживаемый compute_type '{compute_type}' для onnx-asr. "
f"Допустимо: {supported}."
)
return _QUANTIZATION_MAP[compute_type]
class OnnxAsrBackend:
"""Бэкенд транскрипции через onnx-asr (ONNX Runtime)."""
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
def ensure_model_available(
self,
model_name: str,
compute_type: str,
on_status: Callable[[str], None] | None = None,
) -> str:
"""Resolves model alias and returns the onnx-asr model identifier.
onnx-asr downloads models automatically via load_model(),
so this just validates the alias and returns the identifier string.
"""
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
def create_model(
self,
model_path: str,
device: str,
compute_type: str,
cpu_threads: int = 0,
) -> Any:
"""Creates onnx-asr model with VAD.
compute_type маппится в onnx-asr ``quantization`` — это суффикс файла
модели; для unquantized (float32/fp32) нужно None, не строку.
"""
import onnx_asr
actual_compute_type = self.actual_compute_type or compute_type
quantization = _normalize_quantization(actual_compute_type)
model = onnx_asr.load_model(
model=model_path,
quantization=quantization,
)
vad = onnx_asr.load_vad("silero")
self._vad = vad
return model.with_vad(vad)
def transcribe(
self,
model: Any,
file_path: Path,
language: str | None,
on_segment: Callable[[Segment], None] | None = None,
on_status: Callable[[str], None] | None = None,
) -> TranscribeResult:
"""Transcribes audio file using onnx-asr model with VAD.
model: result of create_model() — a SegmentResultsAsrAdapter.
file_path: path to audio/video file (any format supported by faster-whisper decode).
language: language code (e.g. "ru", "en") — only meaningful for multilingual models.
"""
from faster_whisper import decode_audio
_notify(on_status, "Загружаю аудио...")
audio_array = decode_audio(str(file_path), sampling_rate=16000)
duration = len(audio_array) / 16000.0
_notify(on_status, "Транскрибирую (onnx-asr)...")
segments: list[Segment] = []
detected_language = language or "unknown"
for vad_seg in model.recognize(
audio_array, sample_rate=16000, language=language
):
start = max(0.0, vad_seg.start)
end = max(0.0, vad_seg.end)
if end <= start:
continue
seg = Segment(
start=start,
end=end,
text=vad_seg.text,
)
if on_segment is not None:
on_segment(seg)
segments.append(seg)
_notify(
on_status,
f"Транскрибирую (onnx-asr)... [{len(segments)} сегм.]",
)
return TranscribeResult(
segments=segments,
language=detected_language,
language_probability=1.0 if language else 0.0,
duration=duration,
device_used="", # оркестратор проставит
)
def _resolve_model(self, model_name: str) -> str:
"""Resolve alias to onnx-asr model name. Raw names pass through."""
if model_name in MODEL_ALIASES:
return MODEL_ALIASES[model_name]
if "/" in model_name or model_name.count("-") >= 2:
# Looks like a raw onnx-asr name — allow passthrough
return model_name
raise ValueError(
f"Неподдерживаемая модель '{model_name}'. "
f"Доступные алиасы: {SUPPORTED_ALIASES}. "
f"Либо укажите полное имя модели onnx-asr."
)
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)