feat(onnx-asr): реализовать ensure_model_available
- Зачем: - Бэкенд должен резолвить алиасы моделей и сохранять compute_type для дальнейшего использования. - Что: - добавлен конструктор __init__ с полями actual_compute_type, _resolved_model_id, _vad. - метод ensure_model_available валидирует алиас и возвращает идентификатор модели onnx-asr. - написаны 3 теста на резолвинг и сохранение compute_type. - Проверка: - uv run pytest tests/test_onnx_asr.py -v
This commit is contained in:
@@ -2,6 +2,9 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
MODEL_ALIASES: dict[str, str] = {
|
MODEL_ALIASES: dict[str, str] = {
|
||||||
"gigaam-v3": "gigaam-v3-ctc",
|
"gigaam-v3": "gigaam-v3-ctc",
|
||||||
"parakeet-v3": "nemo-parakeet-tdt-0.6b-v3",
|
"parakeet-v3": "nemo-parakeet-tdt-0.6b-v3",
|
||||||
@@ -13,6 +16,26 @@ SUPPORTED_ALIASES = ", ".join(MODEL_ALIASES)
|
|||||||
class OnnxAsrBackend:
|
class OnnxAsrBackend:
|
||||||
"""Бэкенд транскрипции через onnx-asr (ONNX Runtime)."""
|
"""Бэкенд транскрипции через onnx-asr (ONNX Runtime)."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
self.actual_compute_type = compute_type
|
||||||
|
self._resolved_model_id = self._resolve_model(model_name)
|
||||||
|
return self._resolved_model_id
|
||||||
|
|
||||||
def _resolve_model(self, model_name: str) -> str:
|
def _resolve_model(self, model_name: str) -> str:
|
||||||
"""Resolve alias to onnx-asr model name. Raw names pass through."""
|
"""Resolve alias to onnx-asr model name. Raw names pass through."""
|
||||||
if model_name in MODEL_ALIASES:
|
if model_name in MODEL_ALIASES:
|
||||||
|
|||||||
@@ -4,6 +4,24 @@ import pytest
|
|||||||
from local_transcriber.backends.onnx_asr import OnnxAsrBackend, MODEL_ALIASES
|
from local_transcriber.backends.onnx_asr import OnnxAsrBackend, MODEL_ALIASES
|
||||||
|
|
||||||
|
|
||||||
|
class TestEnsureModelAvailable:
|
||||||
|
def test_returns_model_id_for_gigaam(self):
|
||||||
|
backend = OnnxAsrBackend()
|
||||||
|
result = backend.ensure_model_available("gigaam-v3", "int8")
|
||||||
|
assert result == "gigaam-v3-ctc"
|
||||||
|
|
||||||
|
def test_returns_model_id_for_parakeet(self):
|
||||||
|
backend = OnnxAsrBackend()
|
||||||
|
result = backend.ensure_model_available("parakeet-v3", "fp16")
|
||||||
|
assert result == "nemo-parakeet-tdt-0.6b-v3"
|
||||||
|
|
||||||
|
def test_stores_compute_type(self):
|
||||||
|
backend = OnnxAsrBackend()
|
||||||
|
backend.ensure_model_available("gigaam-v3", "float32")
|
||||||
|
assert backend._resolved_model_id == "gigaam-v3-ctc"
|
||||||
|
assert backend.actual_compute_type == "float32"
|
||||||
|
|
||||||
|
|
||||||
class TestModelAliases:
|
class TestModelAliases:
|
||||||
def test_gigaam_v3_resolves(self):
|
def test_gigaam_v3_resolves(self):
|
||||||
backend = OnnxAsrBackend()
|
backend = OnnxAsrBackend()
|
||||||
|
|||||||
Reference in New Issue
Block a user