diff --git a/src/local_transcriber/backends/onnx_asr.py b/src/local_transcriber/backends/onnx_asr.py index 190c6f2..770c412 100644 --- a/src/local_transcriber/backends/onnx_asr.py +++ b/src/local_transcriber/backends/onnx_asr.py @@ -2,6 +2,9 @@ from __future__ import annotations +from collections.abc import Callable +from typing import Any + MODEL_ALIASES: dict[str, str] = { "gigaam-v3": "gigaam-v3-ctc", "parakeet-v3": "nemo-parakeet-tdt-0.6b-v3", @@ -13,6 +16,26 @@ SUPPORTED_ALIASES = ", ".join(MODEL_ALIASES) class OnnxAsrBackend: """Бэкенд транскрипции через 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: """Resolve alias to onnx-asr model name. Raw names pass through.""" if model_name in MODEL_ALIASES: diff --git a/tests/test_onnx_asr.py b/tests/test_onnx_asr.py index 4cceb3a..a1f842d 100644 --- a/tests/test_onnx_asr.py +++ b/tests/test_onnx_asr.py @@ -4,6 +4,24 @@ import pytest 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: def test_gigaam_v3_resolves(self): backend = OnnxAsrBackend()