feat(onnx-asr): добавить скелет бэкенда с разрешением алиасов моделей
- Зачем: - onnx-asr бэкенд должен распознавать короткие алиасы (gigaam-v3, parakeet-v3) и raw-имена (nemo-canary-1b-v2). - Что: - создан файл backends/onnx_asr.py с MODEL_ALIASES и классом OnnxAsrBackend. - метод _resolve_model: преобразует алиасы, пропускает raw-имена, выдаёт ValueError для неизвестных. - написаны 4 теста на разрешение алиасов. - Проверка: - uv run pytest tests/test_onnx_asr.py::TestModelAliases -v
This commit is contained in:
@@ -0,0 +1,27 @@
|
|||||||
|
"""Бэкенд транскрипции на основе onnx-asr (GigaAM, Parakeet, FastConformer)."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
MODEL_ALIASES: dict[str, str] = {
|
||||||
|
"gigaam-v3": "gigaam-v3-ctc",
|
||||||
|
"parakeet-v3": "nemo-parakeet-tdt-0.6b-v3",
|
||||||
|
}
|
||||||
|
|
||||||
|
SUPPORTED_ALIASES = ", ".join(MODEL_ALIASES)
|
||||||
|
|
||||||
|
|
||||||
|
class OnnxAsrBackend:
|
||||||
|
"""Бэкенд транскрипции через onnx-asr (ONNX Runtime)."""
|
||||||
|
|
||||||
|
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."
|
||||||
|
)
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
"""Tests for onnx-asr backend."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from local_transcriber.backends.onnx_asr import OnnxAsrBackend, MODEL_ALIASES
|
||||||
|
|
||||||
|
|
||||||
|
class TestModelAliases:
|
||||||
|
def test_gigaam_v3_resolves(self):
|
||||||
|
backend = OnnxAsrBackend()
|
||||||
|
result = backend._resolve_model("gigaam-v3")
|
||||||
|
assert result == "gigaam-v3-ctc"
|
||||||
|
|
||||||
|
def test_parakeet_v3_resolves(self):
|
||||||
|
backend = OnnxAsrBackend()
|
||||||
|
result = backend._resolve_model("parakeet-v3")
|
||||||
|
assert result == "nemo-parakeet-tdt-0.6b-v3"
|
||||||
|
|
||||||
|
def test_raw_name_passes_through(self):
|
||||||
|
backend = OnnxAsrBackend()
|
||||||
|
result = backend._resolve_model("nemo-canary-1b-v2")
|
||||||
|
assert result == "nemo-canary-1b-v2"
|
||||||
|
|
||||||
|
def test_unknown_alias_raises(self):
|
||||||
|
backend = OnnxAsrBackend()
|
||||||
|
with pytest.raises(ValueError, match="Неподдерживаемая модель"):
|
||||||
|
backend._resolve_model("nonexistent-model")
|
||||||
Reference in New Issue
Block a user